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

Update src/streamlit_app.py

Browse files
Files changed (1) hide show
  1. src/streamlit_app.py +147 -92
src/streamlit_app.py CHANGED
@@ -5,13 +5,17 @@ import requests
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
  # =========================
@@ -46,8 +50,8 @@ tokenizer, model, device = load_model()
46
  # =========================
47
  @st.cache_resource
48
  def get_langsmith():
49
- key = os.getenv("LANGCHAIN_API_KEY") or os.getenv("LANGSMITH_API_KEY")
50
- if key and LangSmithClient:
51
  return LangSmithClient()
52
  return None
53
 
@@ -60,14 +64,17 @@ ls_client = get_langsmith()
60
  # =========================
61
  if "messages" not in st.session_state:
62
  st.session_state.messages = []
 
 
63
 
64
 
65
  # =========================
66
- # HEADER
 
67
  # =========================
68
- col1, col2 = st.columns([1, 6])
69
  with col1:
70
- st.image(LOGO_URL, use_column_width=True)
71
  with col2:
72
  st.markdown("## TeapotAI Chat")
73
  st.caption("Grounded answers with web context")
@@ -124,8 +131,7 @@ def web_search_snippets(query: str):
124
 
125
  snippets = []
126
  for item in data.get("web", {}).get("results", [])[:TOP_K_SEARCH]:
127
- desc = (item.get("description") or "")
128
- desc = desc.replace("<strong>", "").replace("</strong>", "").strip()
129
  if desc:
130
  snippets.append(desc)
131
 
@@ -137,18 +143,16 @@ def web_search_snippets(query: str):
137
  # =========================
138
  def truncate_context(web_ctx: str, local_ctx: str, system: str, question: str) -> str:
139
  ctx = f"{web_ctx}\n\n{local_ctx}".strip()
140
-
141
  base = f"\n{system}\n{question}\n"
142
  base_tokens = tokenizer.encode(base)
143
  budget = MAX_INPUT_TOKENS - len(base_tokens)
144
-
145
  if budget <= 0:
146
  return ""
147
-
148
- ctx_tokens = tokenizer.encode(ctx) if ctx else []
 
149
  if len(ctx_tokens) <= budget:
150
  return ctx
151
-
152
  return tokenizer.decode(ctx_tokens[-budget:], skip_special_tokens=True)
153
 
154
 
@@ -157,44 +161,68 @@ def count_tokens(text: str) -> int:
157
 
158
 
159
  # =========================
160
- # STREAM GENERATION
 
161
  # =========================
162
- def stream_generate(prompt: str):
163
- inputs = tokenizer(prompt, return_tensors="pt").to(device)
164
- streamer = TextIteratorStreamer(tokenizer, skip_special_tokens=True)
165
-
166
- def run():
167
- model.generate(
168
- **inputs,
169
- max_new_tokens=MAX_NEW_TOKENS,
170
- do_sample=False,
171
- num_beams=1,
172
- streamer=streamer,
173
- )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
174
 
175
- threading.Thread(target=run, daemon=True).start()
176
 
177
- acc = ""
178
- for chunk in streamer:
179
- acc += chunk
180
- yield acc
 
 
 
 
 
181
 
182
 
183
  # =========================
184
- # FEEDBACK HANDLER
185
  # =========================
186
  def handle_feedback(idx: int):
187
  val = st.session_state.get(f"fb_{idx}")
188
  st.session_state.messages[idx]["feedback"] = val
189
 
190
  msg = st.session_state.messages[idx]
191
- if ls_client and msg.get("run_id"):
 
 
 
192
  score = 1 if val == "👍" else 0
193
  try:
 
194
  ls_client.create_feedback(
195
- run_id=msg["run_id"],
196
  key="thumb_rating",
197
  score=score,
 
198
  )
199
  except Exception:
200
  pass
@@ -205,26 +233,43 @@ def handle_feedback(idx: int):
205
  # =========================
206
  for i, msg in enumerate(st.session_state.messages):
207
  with st.chat_message(msg["role"]):
 
 
 
 
 
208
  st.markdown(msg["content"])
209
 
210
- if msg["role"] == "assistant":
 
 
 
 
 
 
 
 
 
 
 
 
 
211
  st.caption(
212
- f"{msg['search_time']:.2f}s • {msg['gen_time']:.2f}s • "
213
- f"{msg['tps']:.1f} tok/s • in {msg['input_tokens']} • out {msg['output_tokens']}"
 
 
214
  )
215
 
216
- with st.expander("Inspect context"):
217
- st.code(msg.get("prompt", ""), language="text")
218
-
219
- key = f"fb_{i}"
220
- st.session_state.setdefault(key, msg.get("feedback"))
221
- st.feedback(
222
- "thumbs",
223
- key=key,
224
- disabled=msg.get("feedback") is not None,
225
- on_change=handle_feedback,
226
- args=(i,),
227
- )
228
 
229
 
230
  # =========================
@@ -234,79 +279,89 @@ query = st.chat_input("Ask a question...")
234
 
235
  if query:
236
  st.session_state.messages.append({"role": "user", "content": query})
 
237
  st.rerun()
238
 
239
 
240
  # =========================
241
- # GENERATE
242
  # =========================
243
- if st.session_state.messages and st.session_state.messages[-1]["role"] == "user":
 
 
 
 
244
  question = st.session_state.messages[-1]["content"]
245
 
 
246
  web_ctx, search_time = web_search_snippets(question)
247
 
248
- final_context = truncate_context(
249
- web_ctx,
250
- local_context,
251
- system_prompt,
252
- question,
253
- )
254
-
255
- prompt = f"{final_context}\n{system_prompt}\n{question}\n"
256
  input_tokens = count_tokens(prompt)
257
 
258
- run_id = None
259
- if ls_client:
260
- try:
261
- run = ls_client.create_run(
262
- name="teapot_chat",
263
- run_type="llm",
264
- inputs={"prompt": prompt, "question": question},
265
- )
266
- run_id = run.id
267
- except Exception:
268
- pass
269
-
270
  with st.chat_message("assistant"):
271
  placeholder = st.empty()
 
272
  start = time.perf_counter()
273
- final_text = ""
274
 
275
- for partial in stream_generate(prompt):
276
- final_text = partial
277
- placeholder.markdown(final_text)
 
 
 
 
 
 
 
 
 
278
 
279
  gen_time = time.perf_counter() - start
280
- output_tokens = count_tokens(final_text)
281
  tps = output_tokens / gen_time if gen_time > 0 else 0.0
282
 
283
- st.caption(
284
- f"{search_time:.2f}s • {gen_time:.2f}s • "
285
- f"{tps:.1f} tok/s • in {input_tokens} • out {output_tokens}"
286
- )
287
-
288
- with st.expander("Inspect context"):
289
- st.code(prompt, language="text")
290
-
291
- if ls_client and run_id:
292
- try:
293
- ls_client.update_run(run_id, outputs={"answer": final_text})
294
- except Exception:
295
- pass
 
 
 
 
 
 
296
 
 
297
  st.session_state.messages.append(
298
  {
299
  "role": "assistant",
300
- "content": final_text,
 
 
 
301
  "prompt": prompt,
302
  "search_time": search_time,
303
  "gen_time": gen_time,
304
  "input_tokens": input_tokens,
305
  "output_tokens": output_tokens,
306
  "tps": tps,
307
- "run_id": run_id,
308
  "feedback": None,
309
  }
310
  )
311
 
 
312
  st.rerun()
 
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
  # =========================
 
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
 
 
64
  # =========================
65
  if "messages" not in st.session_state:
66
  st.session_state.messages = []
67
+ if "needs_answer" not in st.session_state:
68
+ st.session_state.needs_answer = False
69
 
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")
 
131
 
132
  snippets = []
133
  for item in data.get("web", {}).get("results", [])[:TOP_K_SEARCH]:
134
+ desc = (item.get("description") or "").replace("<strong>", "").replace("</strong>", "").strip()
 
135
  if desc:
136
  snippets.append(desc)
137
 
 
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
 
 
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
 
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
  # =========================
 
279
 
280
  if query:
281
  st.session_state.messages.append({"role": "user", "content": query})
282
+ st.session_state.needs_answer = True
283
  st.rerun()
284
 
285
 
286
  # =========================
287
+ # GENERATE (once per user message)
288
  # =========================
289
+ if (
290
+ st.session_state.needs_answer
291
+ and st.session_state.messages
292
+ and st.session_state.messages[-1]["role"] == "user"
293
+ ):
294
  question = st.session_state.messages[-1]["content"]
295
 
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
  )
365
 
366
+ st.session_state.needs_answer = False
367
  st.rerun()