diff --git a/api/Procfile b/api/Procfile index 1945201..0d052ab 100644 --- a/api/Procfile +++ b/api/Procfile @@ -1 +1 @@ -web: gunicorn main:app +web: gunicorn main:app --worker-class eventlet --threads 4 diff --git a/api/chat.py b/api/chat.py index 5d5b30e..10fbf1d 100644 --- a/api/chat.py +++ b/api/chat.py @@ -24,7 +24,7 @@ CONTEXT_FRACTION = 0.5 # the (approximate) fraction of num_tokens to use for co ENCODER = tiktoken.get_encoding("cl100k_base") -DEBUG_PRINT = False +DEBUG_PRINT = True # --------------------------------- prompt code -------------------------------- @@ -116,32 +116,44 @@ def construct_prompt(query: str, history: List[Dict[str, str]], context: List[Bl return prompt # ------------------------------- completion code ------------------------------- +import time +import json # returns either (True, reply string, top_k_blocks)) or (False, error message string, None) def talk_to_robot(index, query: str, history: List[Dict[str, str]], k: int = STANDARD_K): - try: # 1. Find the most relevant blocks from the Alignment Research Dataset + yield json.dumps({"state": "loading", "phase": "semantic"}) top_k_blocks = get_top_k_blocks(index, query, k) + yield json.dumps({"state": "loading", "phase": "semantic", 'citations': [{'title': block.title, 'author': block.author, 'date': block.date, 'url': block.url} for block in top_k_blocks]}) # 2. Generate a prompt + yield json.dumps({"state": "loading", "phase": "prompt"}) prompt = construct_prompt(query, history, top_k_blocks) - # 3. Count number of tokens left for completion (-50 for a buffer) max_tokens_completion = NUM_TOKENS - sum([len(ENCODER.encode(message["content"]) + ENCODER.encode(message["role"])) for message in prompt]) - 50 # 4. Answer the user query + yield json.dumps({"state": "loading", "phase": "llm"}) t1 = time.time() - response = openai.ChatCompletion.create( + response = '' + + for chunk in openai.ChatCompletion.create( model=COMPLETIONS_MODEL, messages=prompt, - max_tokens=max_tokens_completion - )["choices"][0]["message"]["content"] + max_tokens=max_tokens_completion, + stream=True + ): + res = chunk["choices"][0]["delta"] + if res is not None and res.get("content") is not None: + response += res["content"] + yield json.dumps({"state": "streaming", "content": res["content"]}) + + t2 = time.time() print("Time to get response: ", t2 - t1) - if DEBUG_PRINT: print('\n' * 10) @@ -155,9 +167,9 @@ def talk_to_robot(index, query: str, history: List[Dict[str, str]], k: int = STA print(" ------------------------------ response: -----------------------------") print(response) - return (True, response, top_k_blocks) + yield json.dumps({"state": "done"}) except Exception as e: print(e) - return (False, "Error: " + str(e), None) + yield json.dumps({"state": "error", "error": str(e)}) diff --git a/api/main.py b/api/main.py index 099c2ca..45f04c9 100644 --- a/api/main.py +++ b/api/main.py @@ -1,4 +1,4 @@ -from flask import Flask, jsonify, request +from flask import Flask, jsonify, request, Response from flask_cors import CORS, cross_origin from get_blocks import get_top_k_blocks from chat import talk_to_robot @@ -34,6 +34,12 @@ app = Flask(__name__) cors = CORS(app) app.config['CORS_HEADERS'] = 'Content-Type' +# ---------------------------------- sse stuff --------------------------------- + +def stream(src): + yield from ('data: ' + '\ndata: '.join(message.splitlines()) + '\n\n' for message in src) + yield 'event: close\n\n' + # ------------------------------- semantic search ------------------------------ @@ -50,20 +56,15 @@ def semantic(): @app.route('/chat', methods=['POST']) @cross_origin() def chat(): - + query = request.json['query'] history = request.json['history'] - is_valid, response, context = talk_to_robot(index, query, history) - - if is_valid: - return jsonify({'response': response, 'citations': [{'title': block.title, 'author': block.author, 'date': block.date, 'url': block.url} for block in context]}) - else: - return jsonify({'error': response}) + return Response(stream(talk_to_robot(index, query, history)), mimetype='text/event-stream') # ------------------------------------------------------------------------------ if __name__ == '__main__': - app.run(debug=True, port=3000) \ No newline at end of file + app.run(debug=True, port=3000) diff --git a/web/src/pages/index.tsx b/web/src/pages/index.tsx index 6f6e397..a770c5d 100644 --- a/web/src/pages/index.tsx +++ b/web/src/pages/index.tsx @@ -23,7 +23,8 @@ type UserEntry = { type AssistantEntry = { role: "assistant"; content: string; - citations: Map; + citations: Citation[]; + base_count: number; // the number to start counting citations at } type ErrorMessage = { @@ -82,35 +83,109 @@ const ShowEntry: React.FC<{entry: Entry}> = ({entry}) => { if (entry.role === "user") { return (

{entry.content}

); } - - // error message - if (entry.role === "error") { - return (

{entry.content}

); - } + + + + +// todo: memoize this if too slow. +const ProcessText: (text: string, base_count: number) => [string, Map] = (text, base_count) => { + + // ---------------------- normalize citation form ---------------------- + + // transform all things that look like [a, b, c] into [a][b][c] + let response = text.replace( + + /\[((?:[a-z]+,\s*)*[a-z]+)\]/g, // identify groups of this form + + (block: string) => block.split(',') + .map((x) => x.trim()) + .join("][") + ) + + // transform all things that look like [(a), (b), (c)] into [(a)][(b)][(c)] + response = response.replace( + + /\[((?:\([a-z]+\),\s*)*\([a-z]+\))\]/g, // identify groups of this form + + (block: string) => block.split(',') + .map((x) => x.trim()) + .join("][") + ) + + // transform all things that look like [(a)] into [a] + response = response.replace( + /\[\(([a-z]+)\)\]/g, + (_match: string, x: string) => `[${x}]` + ) + + // transform all things that look like [ a ] into [a] + response = response.replace( + /\[\s*([a-z]+)\s*\]/g, + (_match: string, x: string) => `[${x}]` + ) + + // -------------- map citations from strings into numbers -------------- + + // figure out what citations are in the response, and map them appropriately + const cite_map = new Map(); + let cite_count = 0; + + // scan a regex for [x] over the response. If x isn't in the map, add it. + const regex = /\[([a-z]+)\]/g; + let match; + let response_copy = "" + while ((match = regex.exec(response)) !== null) { + if (!cite_map.has(match[1]!)) { + cite_map.set(match[1]!, base_count + cite_count++); + } + // replace [x] with [i] + response_copy += response.slice(response_copy.length, match.index) + `[${cite_map.get(match[1]!)! + 1}]`; + } + + response = response_copy + response.slice(response_copy.length); + + return [response, cite_map] +} + + +const ShowAssistantEntry: React.FC<{entry: AssistantEntry}> = ({entry}) => { const in_text_citation_regex = /\[([0-9]+)\]/g; - // system reply + let [response, cite_map] = ProcessText(entry.content, entry.base_count); + + // ----------------- create the ordered citation array ----------------- + + const citations = new Map(); + cite_map.forEach((value, key) => { + const index = key.charCodeAt(0) - 'a'.charCodeAt(0); + if (index >= entry.citations.length) { + console.log("invalid citation index: " + index); + } else { + citations.set(value, entry.citations[index]!); + } + }); + return (
{ // split into paragraphs - entry.content.split("\n").map(paragraph => (

{ + response.split("\n").map(paragraph => (

{ paragraph.split(in_text_citation_regex).map((text, i) => { if (i % 2 === 0) { return text.trim(); } i = parseInt(text) - 1; - if (!entry.citations.has(i)) return `[${text}]`; - const citation = entry.citations.get(i)!; + if (!citations.has(i)) return `[${text}]`; + const citation = citations.get(i)!; return ( ); }) }

)) } -
    +
      { // show citations - Array.from(entry.citations.entries()).map(([i, citation]) => ( + Array.from(citations.entries()).map(([i, citation]) => (
    • @@ -121,29 +196,55 @@ const ShowEntry: React.FC<{entry: Entry}> = ({entry}) => { ); }; + + + + +type State = { + state: "idle"; +} | { + state: "loading"; + phase: "semantic" | "prompt" | "llm"; + citations: Citation[]; +} | { + state: "streaming"; + response: AssistantEntry; +}; + + const Home: NextPage = () => { const [ entries, setEntries ] = useState([]); const [ runningIndex, setRunningIndex ] = useState(0); + const [ loadState, setLoadState ] = useState({state: "idle"}); const search = async ( query: string, setQuery: (query: string) => void, setLoading: (loading: boolean) => void ) => { - + // clear the query box, append to entries + const old_entries = entries; const new_entries: Entry[] = [...old_entries, {role: "user", content: query}]; setEntries(new_entries); setQuery(""); - setLoading(true); + // do SSE on a POST request. + const res = await fetch(API_URL + "/chat", { method: "POST", - headers: { "Content-Type": "application/json", "Allow-Control-Allow-Origin": "*" }, - body: JSON.stringify({query: query, history: + cache: "no-cache", + keepalive: true, + headers: { + "Content-Type": "application/json", + "Accept": "text/event-stream", + "Allow-Control-Allow-Origin": "*" + }, + + body: JSON.stringify({query: query, history: old_entries.filter((entry) => entry.role !== "error") .map((entry) => { return { @@ -151,98 +252,105 @@ const Home: NextPage = () => { "content" : entry.content.trim(), } }) - }) - }) + }), + + }); if (!res.ok) { setLoading(false); - console.log("load failure: " + res.status); + setLoadState({state: "idle"}); + setEntries([...new_entries, {role: "error", content: "POST Error: " + res.status}]); return; } - const data = await res.json(); + // read back the SSE stream - // -------------------------- error checking --------------------------- + const reader = res.body!.getReader(); + var message = ""; + read: while (true) { - if (data.error) { - setEntries([...new_entries, {role: "error", content: data.error}]); - setLoading(false); - return; - } + const {done, value} = await reader.read(); - // ---------------------- normalize citation form ---------------------- + if (done) break; + const chunk = new TextDecoder("utf-8").decode(value); + if (chunk.startsWith("event: close\n")) break; - // transform all things that look like [a, b, c] into [a][b][c] - let response = data.response.replace( + // note: this form isn't even remotely close to optimal in terms of network usage. - /\[((?:[a-z]+,\s*)*[a-z]+)\]/g, // identify groups of this form + for (const line of chunk.split('\n')) { - (block: string) => block.split(',') - .map((x) => x.trim()) - .join("][") - ) + // Most times, it seems that a single read() call will be one SSE "message", + // but I'll do the proper aggregation spec thing in case that's not always true. - // transform all things that look like [(a), (b), (c)] into [(a)][(b)][(c)] - response = response.replace( - - /\[((?:\([a-z]+\),\s*)*\([a-z]+\))\]/g, // identify groups of this form + if (line.startsWith("data: ")) message += line.slice(6); + if (line === "") { + if (message !== "") { + const data = JSON.parse(message); - (block: string) => block.split(',') - .map((x) => x.trim()) - .join("][") - ) + switch (data.state) { - // transform all things that look like [(a)] into [a] - response = response.replace( - /\[\(([a-z]+)\)\]/g, - (_match: string, x: string) => `[${x}]` - ) + case "loading": - // transform all things that look like [ a ] into [a] - response = response.replace( - /\[\s*([a-z]+)\s*\]/g, - (_match: string, x: string) => `[${x}]` - ) + // display loading phases, once citations are available toss them + // into the loading state. - // -------------- map citations from strings into numbers -------------- + setLoadState((s) => { + var citations = s.state === "loading" ? s.citations : []; + if (data.citations !== undefined) { + citations = data.citations; + } + return {state: "loading", phase: data.phase, citations: citations}; + }); - // figure out what citations are in the response, and map them appropriately - const cite_map = new Map(); - let cite_count = runningIndex; + break; - // scan a regex for [x] over the response. If x isn't in the map, add it. - const regex = /\[([a-z]+)\]/g; - let match; - let response_copy = "" - while ((match = regex.exec(response)) !== null) { - if (!cite_map.has(match[1]!)) { - cite_map.set(match[1]!, cite_count++); + case "streaming": + + // incrementally build up the response + + setLoadState((s) => { + const response = s.state === "streaming" ? s.response : + {role: "assistant", + content: "", + citations: s.state === "loading" ? s.citations : [], + base_count: runningIndex + }; + + return {state: "streaming", response: { + role: "assistant", + content: response.content + data.content, + citations: response.citations, + base_count: response.base_count + }}; + }); + break; + + case "done": + + // append the response to the entries, reset to normal + + setLoadState((s) => { + if (s.state === "streaming") { + setEntries([...new_entries, s.response]); + setRunningIndex((i) => (i + ProcessText(s.response.content, 0)[1].size)); + } + return {state: "idle"}; + }); + break read; + + case "error": + setEntries([...new_entries, {role: "error", content: data.error}]); + break read; + + } + } + message = ""; + } } - // replace [x] with [i] - response_copy += response.slice(response_copy.length, match.index) + `[${cite_map.get(match[1]!)! + 1}]`; } - setRunningIndex(cite_count); - - response = response_copy + response.slice(response_copy.length); - - // ----------------- create the ordered citation array ----------------- - - const citations = new Map(); - cite_map.forEach((value, key) => { - const index = key.charCodeAt(0) - 'a'.charCodeAt(0); - if (index >= data.citations.length) { - console.log("invalid citation index: " + index); - } else { - citations.set(value, data.citations[index]); - } - }); - - setEntries([...new_entries, {role: "assistant", - content: response, - citations: citations}]); - setLoading(false); + setLoadState({state: "idle"}); }; @@ -254,13 +362,41 @@ const Home: NextPage = () => {
        - {entries.map((entry, i) => ( -
      • - -
      • - ))} + {entries.map((entry, i) => { + if (entry.role === "user") { + return
      • +

        {entry.content}

        +
      • + } + if (entry.role === "error") { + return
      • +

        {entry.content}

        +
      • + } + if (entry.role === "assistant") { + return
      • + +
      • + } + return <> + })} + + + + {(() => { + if (loadState.state === "loading") { + switch (loadState.phase) { + case "semantic": return

        Loading: Performing semantic search...

        ; + case "prompt": return

        Loading: Creating prompt...

        ; + case "llm": return

        Loading: Waiting for LLM...

        ; + } + } else if (loadState.state === "streaming") { + return ; + } + return <>; + })()} +
      -
      ); diff --git a/web/src/searchbox.tsx b/web/src/searchbox.tsx index d21dedd..9b4e7dd 100644 --- a/web/src/searchbox.tsx +++ b/web/src/searchbox.tsx @@ -20,7 +20,7 @@ const SearchBox: React.FC<{search: ( if (!loading) inputRef.current?.focus(); }, [loading]); - if (loading) return

      loading...

      ; + if (loading) return <>; return (<>
      { e.preventDefault();