zakerytclarke commited on
Commit
1cff0af
·
verified ·
1 Parent(s): afed0d7

Update src/streamlit_app.py

Browse files
Files changed (1) hide show
  1. src/streamlit_app.py +157 -154
src/streamlit_app.py CHANGED
@@ -5,17 +5,13 @@ import requests
5
 
6
  import streamlit as st
7
  import torch
8
- from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
9
 
10
- # Optional LangSmith (trace + feedback)
11
  try:
12
  from langsmith import Client as LangSmithClient
13
- from langsmith import traceable
14
- from langsmith.run_helpers import get_current_run_tree
15
  except Exception:
16
  LangSmithClient = None
17
- traceable = None
18
- get_current_run_tree = None
19
 
20
 
21
  # =========================
@@ -27,6 +23,9 @@ MAX_NEW_TOKENS = 192
27
  TOP_K_SEARCH = 3
28
  LOGO_URL = "https://teapotai.com/assets/logo.gif"
29
 
 
 
 
30
  st.set_page_config(page_title="TeapotAI Chat", page_icon="🫖", layout="centered")
31
 
32
 
@@ -50,8 +49,8 @@ tokenizer, model, device = load_model()
50
  # =========================
51
  @st.cache_resource
52
  def get_langsmith():
53
- key = os.getenv("LANGCHAIN_API_KEY") or os.getenv("LANGSMITH_API_KEY") or os.getenv("LANGCHAIN_TRACING_V2")
54
- if (os.getenv("LANGCHAIN_API_KEY") or os.getenv("LANGSMITH_API_KEY")) and LangSmithClient:
55
  return LangSmithClient()
56
  return None
57
 
@@ -70,11 +69,10 @@ if "needs_answer" not in st.session_state:
70
 
71
  # =========================
72
  # HEADER (prevent logo flash)
73
- # Use a fixed pixel width to avoid layout shift / big flash.
74
  # =========================
75
  col1, col2 = st.columns([1, 7], vertical_alignment="center")
76
  with col1:
77
- st.image(LOGO_URL, width=56) # fixed width prevents "flash huge"
78
  with col2:
79
  st.markdown("## TeapotAI Chat")
80
  st.caption("Grounded answers with web context")
@@ -104,6 +102,14 @@ with st.sidebar:
104
  placeholder="Extra context appended after web snippets…",
105
  )
106
 
 
 
 
 
 
 
 
 
107
 
108
  # =========================
109
  # WEB SEARCH (ALWAYS ON)
@@ -139,93 +145,101 @@ def web_search_snippets(query: str):
139
 
140
 
141
  # =========================
142
- # CONTEXT TRUNCATION (TAIL)
143
  # =========================
144
- def truncate_context(web_ctx: str, local_ctx: str, system: str, question: str) -> str:
145
- ctx = f"{web_ctx}\n\n{local_ctx}".strip()
146
- base = f"\n{system}\n{question}\n"
147
- base_tokens = tokenizer.encode(base)
148
- budget = MAX_INPUT_TOKENS - len(base_tokens)
149
- if budget <= 0:
150
- return ""
151
- if not ctx:
 
152
  return ""
153
- ctx_tokens = tokenizer.encode(ctx)
154
- if len(ctx_tokens) <= budget:
155
- return ctx
156
- return tokenizer.decode(ctx_tokens[-budget:], skip_special_tokens=True)
157
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
158
 
159
- def count_tokens(text: str) -> int:
160
- return len(tokenizer.encode(text)) if text else 0
 
 
 
 
 
 
 
 
 
161
 
162
 
163
  # =========================
164
- # LANGSMITH-TRACED ANSWER FUNCTION
165
- # (signature exactly: context, system_prompt, question -> answer)
166
  # =========================
167
- if traceable:
168
- @traceable(name="teapot_answer")
169
- def traced_answer(context: str, system_prompt: str, question: str) -> str:
170
- prompt = f"{context}\n{system_prompt}\n{question}\n"
171
- inputs = tokenizer(prompt, return_tensors="pt").to(device)
172
- with torch.no_grad():
173
- out = model.generate(
174
- **inputs,
175
- max_new_tokens=MAX_NEW_TOKENS,
176
- do_sample=False,
177
- num_beams=1,
178
- )
179
- text = tokenizer.decode(out[0], skip_special_tokens=True)
180
- return text
181
- else:
182
- def traced_answer(context: str, system_prompt: str, question: str) -> str:
183
- prompt = f"{context}\n{system_prompt}\n{question}\n"
184
- inputs = tokenizer(prompt, return_tensors="pt").to(device)
185
- with torch.no_grad():
186
- out = model.generate(
187
- **inputs,
188
- max_new_tokens=MAX_NEW_TOKENS,
189
- do_sample=False,
190
- num_beams=1,
191
- )
192
- return tokenizer.decode(out[0], skip_special_tokens=True)
193
 
 
194
 
195
- def get_trace_id_if_available() -> str | None:
196
- # Works when running inside a @traceable function call
197
- if not get_current_run_tree:
198
- return None
199
- try:
200
- run = get_current_run_tree()
201
- return str(run.id) if run and getattr(run, "id", None) else None
202
- except Exception:
203
- return None
204
 
205
 
206
  # =========================
207
- # FEEDBACK HANDLER (attached to trace_id)
208
  # =========================
209
  def handle_feedback(idx: int):
210
  val = st.session_state.get(f"fb_{idx}")
211
  st.session_state.messages[idx]["feedback"] = val
212
 
213
- msg = st.session_state.messages[idx]
214
- trace_id = msg.get("trace_id")
215
-
216
- # Attach feedback to this traced run
217
- if ls_client and trace_id:
218
- score = 1 if val == "👍" else 0
219
- try:
220
- # LangSmith SDK supports trace_id= for feedback association
221
- ls_client.create_feedback(
222
- trace_id=trace_id,
223
- key="thumb_rating",
224
- score=score,
225
- comment="thumbs_up" if score else "thumbs_down",
226
- )
227
- except Exception:
228
- pass
229
 
230
 
231
  # =========================
@@ -233,43 +247,35 @@ def handle_feedback(idx: int):
233
  # =========================
234
  for i, msg in enumerate(st.session_state.messages):
235
  with st.chat_message(msg["role"]):
236
- if msg["role"] == "user":
237
- st.markdown(msg["content"])
238
- continue
239
-
240
- # Assistant
241
  st.markdown(msg["content"])
242
 
243
- # Info icon popover with full prompt/context
244
- # (st.popover is stable in your Streamlit range; no rerun on open/close)
245
- c1, c2 = st.columns([1, 12], vertical_alignment="center")
246
- with c1:
247
- with st.popover("ℹ️", help="Inspect"):
248
- st.markdown("**Context**")
249
- st.code(msg.get("context", ""), language="text")
250
- st.markdown("**System**")
251
- st.code(msg.get("system_prompt", ""), language="text")
252
- st.markdown("**Question**")
253
- st.code(msg.get("question", ""), language="text")
254
- st.markdown("**Prompt**")
255
- st.code(msg.get("prompt", ""), language="text")
256
- with c2:
257
- st.caption(
258
- f"🔎 {msg['search_time']:.2f}s "
259
- f"🧠 {msg['gen_time']:.2f}s "
260
- f"⚡ {msg['tps']:.1f} tok/s "
261
- f"🧾 in {msg['input_tokens']} • out {msg['output_tokens']}"
262
- )
263
-
264
- key = f"fb_{i}"
265
- st.session_state.setdefault(key, msg.get("feedback"))
266
- st.feedback(
267
- "thumbs",
268
- key=key,
269
- disabled=msg.get("feedback") is not None,
270
- on_change=handle_feedback,
271
- args=(i,),
272
- )
273
 
274
 
275
  # =========================
@@ -296,69 +302,66 @@ if (
296
  # Web search
297
  web_ctx, search_time = web_search_snippets(question)
298
 
299
- # Context + truncation
300
- context = truncate_context(web_ctx, local_context, system_prompt, question)
301
- prompt = f"{context}\n{system_prompt}\n{question}\n"
 
 
 
 
 
 
 
 
 
 
302
  input_tokens = count_tokens(prompt)
303
 
304
- # Run traced answer (returns answer; trace_id obtained from current run tree)
305
  with st.chat_message("assistant"):
306
  placeholder = st.empty()
307
-
308
  start = time.perf_counter()
 
309
 
310
- # Generate full answer first (traced), then "stream" it to UI quickly.
311
- # This keeps LangSmith tracing simple/reliable while still giving a streaming UX.
312
- answer = traced_answer(context, system_prompt, question)
313
- trace_id = get_trace_id_if_available()
314
-
315
- # Typewriter-ish stream (fast, looks normal)
316
- buf = ""
317
- for ch in answer:
318
- buf += ch
319
- placeholder.markdown(buf)
320
- # small delay; tune if you want faster/slower
321
- time.sleep(0.002)
322
 
323
  gen_time = time.perf_counter() - start
324
- output_tokens = count_tokens(answer)
325
  tps = output_tokens / gen_time if gen_time > 0 else 0.0
326
 
327
- # Metrics + info popover for this live message
328
- c1, c2 = st.columns([1, 12], vertical_alignment="center")
329
- with c1:
330
- with st.popover("ℹ️", help="Inspect"):
331
- st.markdown("**Context**")
332
- st.code(context, language="text")
333
- st.markdown("**System**")
334
- st.code(system_prompt, language="text")
335
- st.markdown("**Question**")
336
- st.code(question, language="text")
337
- st.markdown("**Prompt**")
338
  st.code(prompt, language="text")
339
- with c2:
 
 
 
 
 
340
  st.caption(
341
- f"🔎 {search_time:.2f}s "
342
- f"🧠 {gen_time:.2f}s "
343
- f"⚡ {tps:.1f} tok/s "
344
- f"🧾 in {input_tokens} • out {output_tokens}"
345
  )
346
 
347
- # Persist assistant message
348
  st.session_state.messages.append(
349
  {
350
  "role": "assistant",
351
- "content": answer,
352
- "context": context,
353
- "system_prompt": system_prompt,
354
- "question": question,
355
  "prompt": prompt,
356
  "search_time": search_time,
357
  "gen_time": gen_time,
358
  "input_tokens": input_tokens,
359
  "output_tokens": output_tokens,
360
  "tps": tps,
361
- "trace_id": trace_id,
362
  "feedback": None,
363
  }
364
  )
 
5
 
6
  import streamlit as st
7
  import torch
8
+ from transformers import AutoTokenizer, AutoModelForSeq2SeqLM, TextIteratorStreamer
9
 
10
+ # Optional LangSmith
11
  try:
12
  from langsmith import Client as LangSmithClient
 
 
13
  except Exception:
14
  LangSmithClient = None
 
 
15
 
16
 
17
  # =========================
 
23
  TOP_K_SEARCH = 3
24
  LOGO_URL = "https://teapotai.com/assets/logo.gif"
25
 
26
+ # How many (user,assistant) pairs to include in the prompt by default
27
+ MAX_TURNS_IN_PROMPT = 6
28
+
29
  st.set_page_config(page_title="TeapotAI Chat", page_icon="🫖", layout="centered")
30
 
31
 
 
49
  # =========================
50
  @st.cache_resource
51
  def get_langsmith():
52
+ key = os.getenv("LANGCHAIN_API_KEY") or os.getenv("LANGSMITH_API_KEY")
53
+ if key and LangSmithClient:
54
  return LangSmithClient()
55
  return None
56
 
 
69
 
70
  # =========================
71
  # HEADER (prevent logo flash)
 
72
  # =========================
73
  col1, col2 = st.columns([1, 7], vertical_alignment="center")
74
  with col1:
75
+ st.image(LOGO_URL, width=56)
76
  with col2:
77
  st.markdown("## TeapotAI Chat")
78
  st.caption("Grounded answers with web context")
 
102
  placeholder="Extra context appended after web snippets…",
103
  )
104
 
105
+ max_turns = st.slider(
106
+ "Conversation turns in prompt",
107
+ min_value=0,
108
+ max_value=12,
109
+ value=MAX_TURNS_IN_PROMPT,
110
+ help="How many recent (user, assistant) pairs to include in the prompt.",
111
+ )
112
+
113
 
114
  # =========================
115
  # WEB SEARCH (ALWAYS ON)
 
145
 
146
 
147
  # =========================
148
+ # CONTEXT + PROMPT BUILDING
149
  # =========================
150
+ def count_tokens(text: str) -> int:
151
+ return len(tokenizer.encode(text)) if text else 0
152
+
153
+
154
+ def build_conversation(messages, turns: int) -> str:
155
+ """
156
+ Build a compact transcript from the last `turns` (user,assistant) pairs.
157
+ """
158
+ if turns <= 0:
159
  return ""
 
 
 
 
160
 
161
+ # Collect last 2*turns messages ending at the most recent message
162
+ # Keep only user/assistant roles.
163
+ filtered = [m for m in messages if m.get("role") in ("user", "assistant")]
164
+
165
+ # Take tail, but ensure we start on a user message if possible
166
+ tail = filtered[-(2 * turns) :]
167
+ # If first is assistant, drop it (misaligned pair)
168
+ if tail and tail[0]["role"] == "assistant":
169
+ tail = tail[1:]
170
+
171
+ lines = []
172
+ for m in tail:
173
+ role = "User" if m["role"] == "user" else "Assistant"
174
+ content = (m.get("content") or "").strip()
175
+ if content:
176
+ lines.append(f"{role}: {content}")
177
+ return "\n".join(lines).strip()
178
+
179
+
180
+ def truncate_to_token_budget(full_prompt: str, max_tokens: int) -> str:
181
+ ids = tokenizer.encode(full_prompt)
182
+ if len(ids) <= max_tokens:
183
+ return full_prompt
184
+ # Tail truncate to keep the most recent instruction + question
185
+ ids = ids[-max_tokens:]
186
+ return tokenizer.decode(ids, skip_special_tokens=True)
187
+
188
+
189
+ def build_prompt(web_ctx: str, local_ctx: str, system: str, convo: str, question: str) -> str:
190
+ # Order matters: context first, then system, then convo, then question.
191
+ parts = []
192
+ ctx = f"{web_ctx}\n\n{local_ctx}".strip()
193
+ if ctx:
194
+ parts.append("Context:\n" + ctx)
195
 
196
+ parts.append("System:\n" + system.strip())
197
+
198
+ if convo:
199
+ parts.append("Conversation:\n" + convo)
200
+
201
+ parts.append("User:\n" + question.strip())
202
+ parts.append("Assistant:\n") # encourages continuation style
203
+
204
+ raw = "\n\n".join(parts).strip() + "\n"
205
+ # Enforce input budget at token level
206
+ return truncate_to_token_budget(raw, MAX_INPUT_TOKENS)
207
 
208
 
209
  # =========================
210
+ # STREAM GENERATION
 
211
  # =========================
212
+ def stream_generate(prompt: str):
213
+ inputs = tokenizer(prompt, return_tensors="pt").to(device)
214
+ streamer = TextIteratorStreamer(tokenizer, skip_special_tokens=True)
215
+
216
+ def run():
217
+ model.generate(
218
+ **inputs,
219
+ max_new_tokens=MAX_NEW_TOKENS,
220
+ do_sample=False,
221
+ num_beams=1,
222
+ streamer=streamer,
223
+ )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
224
 
225
+ threading.Thread(target=run, daemon=True).start()
226
 
227
+ acc = ""
228
+ for chunk in streamer:
229
+ acc += chunk
230
+ yield acc
 
 
 
 
 
231
 
232
 
233
  # =========================
234
+ # FEEDBACK
235
  # =========================
236
  def handle_feedback(idx: int):
237
  val = st.session_state.get(f"fb_{idx}")
238
  st.session_state.messages[idx]["feedback"] = val
239
 
240
+ # If you later add LangSmith run_ids per message, hook it here.
241
+ # (Keeping this simple/stable like your previous version.)
242
+ # if ls_client and st.session_state.messages[idx].get("run_id"): ...
 
 
 
 
 
 
 
 
 
 
 
 
 
243
 
244
 
245
  # =========================
 
247
  # =========================
248
  for i, msg in enumerate(st.session_state.messages):
249
  with st.chat_message(msg["role"]):
 
 
 
 
 
250
  st.markdown(msg["content"])
251
 
252
+ if msg["role"] == "assistant":
253
+ # Inline row: info popover + thumbs + metrics
254
+ c_info, c_fb, c_metrics = st.columns([1.2, 1.4, 10], vertical_alignment="center")
255
+
256
+ with c_info:
257
+ with st.popover("ℹ️", help="Inspect prompt"):
258
+ st.markdown("**Prompt sent to model**")
259
+ st.code(msg.get("prompt", ""), language="text")
260
+
261
+ with c_fb:
262
+ key = f"fb_{i}"
263
+ st.session_state.setdefault(key, msg.get("feedback"))
264
+ st.feedback(
265
+ "thumbs",
266
+ key=key,
267
+ disabled=msg.get("feedback") is not None,
268
+ on_change=handle_feedback,
269
+ args=(i,),
270
+ )
271
+
272
+ with c_metrics:
273
+ st.caption(
274
+ f"🔎 {msg['search_time']:.2f}s "
275
+ f"• 🧠 {msg['gen_time']:.2f}s "
276
+ f"• ⚡ {msg['tps']:.1f} tok/s "
277
+ f"• 🧾 in {msg['input_tokens']} • out {msg['output_tokens']}"
278
+ )
 
 
 
279
 
280
 
281
  # =========================
 
302
  # Web search
303
  web_ctx, search_time = web_search_snippets(question)
304
 
305
+ # Conversation transcript (from prior messages, excluding current user msg is fine either way;
306
+ # keeping it includes the last user msg too, but we also add question explicitly.)
307
+ convo = build_conversation(st.session_state.messages[:-1], turns=max_turns)
308
+
309
+ # Prompt
310
+ prompt = build_prompt(
311
+ web_ctx=web_ctx,
312
+ local_ctx=local_context,
313
+ system=system_prompt,
314
+ convo=convo,
315
+ question=question,
316
+ )
317
+
318
  input_tokens = count_tokens(prompt)
319
 
320
+ # Stream normally
321
  with st.chat_message("assistant"):
322
  placeholder = st.empty()
 
323
  start = time.perf_counter()
324
+ final_text = ""
325
 
326
+ for partial in stream_generate(prompt):
327
+ final_text = partial
328
+ placeholder.markdown(final_text)
 
 
 
 
 
 
 
 
 
329
 
330
  gen_time = time.perf_counter() - start
331
+ output_tokens = count_tokens(final_text)
332
  tps = output_tokens / gen_time if gen_time > 0 else 0.0
333
 
334
+ # Inline row under the live message
335
+ c_info, c_fb, c_metrics = st.columns([1.2, 1.4, 10], vertical_alignment="center")
336
+
337
+ with c_info:
338
+ with st.popover("ℹ️", help="Inspect prompt"):
339
+ st.markdown("**Prompt sent to model**")
 
 
 
 
 
340
  st.code(prompt, language="text")
341
+
342
+ # For the live message, we don't have a saved index yet; show disabled thumbs placeholder
343
+ with c_fb:
344
+ st.feedback("thumbs", key="fb_live", disabled=True)
345
+
346
+ with c_metrics:
347
  st.caption(
348
+ f"🔎 {search_time:.2f}s "
349
+ f"• 🧠 {gen_time:.2f}s "
350
+ f"• ⚡ {tps:.1f} tok/s "
351
+ f"• 🧾 in {input_tokens} • out {output_tokens}"
352
  )
353
 
354
+ # Persist assistant message (so feedback attaches properly after rerun)
355
  st.session_state.messages.append(
356
  {
357
  "role": "assistant",
358
+ "content": final_text,
 
 
 
359
  "prompt": prompt,
360
  "search_time": search_time,
361
  "gen_time": gen_time,
362
  "input_tokens": input_tokens,
363
  "output_tokens": output_tokens,
364
  "tps": tps,
 
365
  "feedback": None,
366
  }
367
  )