From 382430da323c998d8a8a11d8cf25024f8afe6a91 Mon Sep 17 00:00:00 2001 From: Daniel O'Connell Date: Fri, 29 Sep 2023 17:08:26 +0200 Subject: [PATCH] Session ids --- api/main.py | 16 +++++++++------- api/src/stampy_chat/chat.py | 12 +++++++++--- api/src/stampy_chat/logging.py | 4 ++-- web/src/hooks/useSearch.ts | 14 +++++++++----- web/src/pages/index.tsx | 6 +++++- 5 files changed, 34 insertions(+), 18 deletions(-) diff --git a/api/main.py b/api/main.py index f18ef68..fc9f9ed 100644 --- a/api/main.py +++ b/api/main.py @@ -1,9 +1,10 @@ -from flask import Flask, jsonify, request, Response -from flask_cors import CORS, cross_origin -from urllib.parse import unquote import dataclasses import json import re +from urllib.parse import unquote + +from flask import Flask, jsonify, request, Response +from flask_cors import CORS, cross_origin from stampy_chat import logging from stampy_chat.env import PINECONE_INDEX, FLASK_PORT @@ -42,11 +43,12 @@ def semantic(): @cross_origin() def chat(): - query = request.json['query'] - mode = request.json['mode'] - history = request.json['history'] + query = request.json.get('query') + mode = request.json.get('mode', 'default') + session_id = request.json.get('sessionId') + history = request.json.get('history', []) - return Response(stream(talk_to_robot(PINECONE_INDEX, query, mode, history)), mimetype='text/event-stream') + return Response(stream(talk_to_robot(PINECONE_INDEX, query, mode, history, session_id)), mimetype='text/event-stream') # ------------- simplified non-streaming chat for internal testing ------------- diff --git a/api/src/stampy_chat/chat.py b/api/src/stampy_chat/chat.py index 59a4845..9656028 100644 --- a/api/src/stampy_chat/chat.py +++ b/api/src/stampy_chat/chat.py @@ -156,13 +156,19 @@ def remaining_tokens(prompt: Prompt): return NUM_TOKENS - used_tokens - TOKENS_BUFFER -def talk_to_robot_internal(index, query: str, mode: str, history: Prompt, k: int = STANDARD_K): +def talk_to_robot_internal(index, query: str, mode: str, history: Prompt, session_id: str, k: int = STANDARD_K): try: # 1. Find the most relevant blocks from the Alignment Research Dataset yield {"state": "loading", "phase": "semantic"} top_k_blocks = get_top_k_blocks(index, query, k) - yield {"state": "loading", "phase": "semantic", 'citations': [{'title': block.title, 'author': block.authors, 'date': block.date, 'url': block.url} for block in top_k_blocks]} + yield { + "state": "loading", "phase": "semantic", + "citations": [ + {'title': block.title, 'author': block.authors, 'date': block.date, 'url': block.url} + for block in top_k_blocks + ] + } # 2. Generate a prompt yield {"state": "loading", "phase": "prompt"} @@ -205,7 +211,7 @@ def talk_to_robot_internal(index, query: str, mode: str, history: Prompt, k: int logger.debug(' ------------------------------ response: -----------------------------') logger.debug(response) - logger.interaction(query, response, history, prompt, top_k_blocks) + logger.interaction(session_id, query, response, history, prompt, top_k_blocks) # yield done state, possibly with followup questions fin_json = {'state': 'done'} diff --git a/api/src/stampy_chat/logging.py b/api/src/stampy_chat/logging.py index c58c827..ab126a1 100644 --- a/api/src/stampy_chat/logging.py +++ b/api/src/stampy_chat/logging.py @@ -41,13 +41,13 @@ class ChatLogger(Logger): def is_debug(self): return self.isEnabledFor(DEBUG) - def interaction(self, query, response, history, prompt, blocks): + def interaction(self, session_id, query, response, history, prompt, blocks): prompt = [i for i in prompt if i.get('role') == 'system'] prompt = prompt[0].get('content') if prompt else None self.item_adder.add( Interaction( - # session_id=session_id, + session_id=session_id, interaction_no=len([i for i in history if i.get('role') == 'user']), query=query, prompt=prompt, diff --git a/web/src/hooks/useSearch.ts b/web/src/hooks/useSearch.ts index 8cceb0f..44e48ee 100644 --- a/web/src/hooks/useSearch.ts +++ b/web/src/hooks/useSearch.ts @@ -97,6 +97,7 @@ export const extractAnswer = async ( }; const fetchLLM = async ( + sessionId: string, query: string, mode: string, history: HistoryEntry[] @@ -110,7 +111,7 @@ const fetchLLM = async ( Accept: "text/event-stream", }, - body: JSON.stringify({ query, mode, history }), + body: JSON.stringify({ sessionId, query, mode, history }), }); export const queryLLM = async ( @@ -118,10 +119,11 @@ export const queryLLM = async ( mode: string, history: HistoryEntry[], baseReferencesIndex: number, - setCurrent: (e?: CurrentSearch) => void + setCurrent: (e?: CurrentSearch) => void, + sessionId: string ): Promise => { // do SSE on a POST request. - const res = await fetchLLM(query, mode, history); + const res = await fetchLLM(sessionId, query, mode, history); if (!res.ok) { return { result: { role: "error", content: "POST Error: " + res.status } }; @@ -191,7 +193,8 @@ export const runSearch = async ( mode: string, baseReferencesIndex: number, entries: Entry[], - setCurrent: (c: CurrentSearch) => void + setCurrent: (c: CurrentSearch) => void, + sessionId: string ): SearchResult => { if (query_source === "search") { const history = entries @@ -206,7 +209,8 @@ export const runSearch = async ( mode, history, baseReferencesIndex, - setCurrent + setCurrent, + sessionId ); } else { // ----------------- HUMAN AUTHORED CONTENT RETRIEVAL ------------------ diff --git a/web/src/pages/index.tsx b/web/src/pages/index.tsx index 0c68982..5cb65ae 100644 --- a/web/src/pages/index.tsx +++ b/web/src/pages/index.tsx @@ -45,6 +45,7 @@ const Home: NextPage = () => { const [entries, setEntries] = useState([]); const [runningIndex, setRunningIndex] = useState(0); const [current, setCurrent] = useState(); + const [sessionId, setSessionId] = useState() // [state, ready to save to localstorage] const [mode, setMode] = useState<[Mode, boolean]>(["default", false]); @@ -52,12 +53,14 @@ const Home: NextPage = () => { // store mode in localstorage useEffect(() => { if (mode[1]) localStorage.setItem("chat_mode", mode[0]); + }, [mode]); // initial load useEffect(() => { const mode = localStorage.getItem("chat_mode") as Mode || "default"; setMode([mode, true]); + setSessionId(crypto.randomUUID()); }, []); const updateCurrent = (current: CurrentSearch) => { @@ -88,7 +91,8 @@ const Home: NextPage = () => { mode[0], runningIndex, entries, - updateCurrent + updateCurrent, + sessionId, ); setCurrent(undefined);