zakerytclarke commited on
Commit
1066a1b
·
verified ·
1 Parent(s): f1d79a3

Update src/streamlit_app.py

Browse files
Files changed (1) hide show
  1. src/streamlit_app.py +164 -124
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
 
@@ -65,17 +69,15 @@ if "needs_answer" not in st.session_state:
65
 
66
 
67
  # =========================
68
- # HEADER (reduce logo flash)
69
- # - fixed width prevents "giant image then shrink"
70
- # - container keeps layout stable
71
  # =========================
72
- with st.container():
73
- col1, col2 = st.columns([1, 7], vertical_alignment="center")
74
- with col1:
75
- st.image(LOGO_URL, width=56) # fixed width = stable, no huge flash
76
- with col2:
77
- st.markdown("## TeapotAI Chat")
78
- st.caption("Grounded answers with web context")
79
 
80
 
81
  # =========================
@@ -137,94 +139,93 @@ def web_search_snippets(query: str):
137
 
138
 
139
  # =========================
140
- # PROMPT BUILDING (NO CONVERSATION HISTORY)
141
  # =========================
142
- def count_tokens(text: str) -> int:
143
- return len(tokenizer.encode(text)) if text else 0
144
-
145
-
146
- def truncate_to_token_budget(text: str, max_tokens: int) -> str:
147
- ids = tokenizer.encode(text)
148
- if len(ids) <= max_tokens:
149
- return text
150
- ids = ids[-max_tokens:] # tail truncate
151
- return tokenizer.decode(ids, skip_special_tokens=True)
 
 
 
152
 
153
 
154
- def build_prompt(web_ctx: str, local_ctx: str, system: str, question: str) -> str:
155
- ctx = f"{web_ctx}\n\n{local_ctx}".strip()
156
- parts = []
157
- if ctx:
158
- parts.append(ctx)
159
- parts.append(system.strip())
160
- parts.append(question.strip())
161
- raw = "\n\n".join(parts).strip() + "\n"
162
- return truncate_to_token_budget(raw, MAX_INPUT_TOKENS)
163
 
164
 
165
  # =========================
166
- # STREAM GENERATION
 
167
  # =========================
168
- def stream_generate(prompt: str):
169
- inputs = tokenizer(prompt, return_tensors="pt").to(device)
170
- streamer = TextIteratorStreamer(tokenizer, skip_special_tokens=True)
171
-
172
- def run():
173
- model.generate(
174
- **inputs,
175
- max_new_tokens=MAX_NEW_TOKENS,
176
- do_sample=False,
177
- num_beams=1,
178
- streamer=streamer,
179
- )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
180
 
181
- threading.Thread(target=run, daemon=True).start()
182
 
183
- acc = ""
184
- for chunk in streamer:
185
- acc += chunk
186
- yield acc
 
 
 
 
 
187
 
188
 
189
  # =========================
190
- # FEEDBACK
191
  # =========================
192
  def handle_feedback(idx: int):
193
  val = st.session_state.get(f"fb_{idx}")
194
  st.session_state.messages[idx]["feedback"] = val
195
 
196
- # If you want to wire LangSmith feedback to a run later, store run_id per message and use it here.
197
- # For now we keep it stable and local like the earlier version.
198
-
199
-
200
- def render_inline_controls(msg: dict, feedback_key: str, feedback_disabled: bool, feedback_idx: int | None):
201
- """
202
- Inline row under assistant message:
203
- ℹ️ popover (full prompt), thumbs, metrics
204
- """
205
- c_info, c_fb, c_metrics = st.columns([1.1, 1.7, 10], vertical_alignment="center")
206
-
207
- with c_info:
208
- with st.popover("ℹ️", help="Inspect"):
209
- st.markdown("**Prompt (sent to model)**")
210
- st.code(msg.get("prompt", ""), language="text")
211
-
212
- with c_fb:
213
- st.feedback(
214
- "thumbs",
215
- key=feedback_key,
216
- disabled=feedback_disabled,
217
- on_change=(handle_feedback if (not feedback_disabled and feedback_idx is not None) else None),
218
- args=((feedback_idx,) if (not feedback_disabled and feedback_idx is not None) else None),
219
- )
220
-
221
- with c_metrics:
222
- st.caption(
223
- f"🔎 {msg.get('search_time', 0.0):.2f}s "
224
- f"• 🧠 {msg.get('gen_time', 0.0):.2f}s "
225
- f"• ⚡ {msg.get('tps', 0.0):.1f} tok/s "
226
- f"• 🧾 in {msg.get('input_tokens', 0)} • out {msg.get('output_tokens', 0)}"
227
- )
228
 
229
 
230
  # =========================
@@ -232,20 +233,44 @@ def render_inline_controls(msg: dict, feedback_key: str, feedback_disabled: bool
232
  # =========================
233
  for i, msg in enumerate(st.session_state.messages):
234
  with st.chat_message(msg["role"]):
235
- st.markdown(msg["content"])
 
 
236
 
237
- if msg["role"] == "assistant":
238
- # Ensure feedback state exists
239
- k = f"fb_{i}"
240
- st.session_state.setdefault(k, msg.get("feedback"))
241
 
242
- render_inline_controls(
243
- msg=msg,
244
- feedback_key=k,
245
- feedback_disabled=(msg.get("feedback") is not None),
246
- feedback_idx=i,
 
 
 
 
 
 
 
 
 
 
 
 
 
 
247
  )
248
 
 
 
 
 
 
 
 
 
 
 
249
 
250
  # =========================
251
  # INPUT
@@ -268,57 +293,72 @@ if (
268
  ):
269
  question = st.session_state.messages[-1]["content"]
270
 
 
271
  web_ctx, search_time = web_search_snippets(question)
272
 
273
- prompt = build_prompt(
274
- web_ctx=web_ctx,
275
- local_ctx=local_context,
276
- system=system_prompt,
277
- question=question,
278
- )
279
  input_tokens = count_tokens(prompt)
280
 
 
281
  with st.chat_message("assistant"):
282
  placeholder = st.empty()
 
283
  start = time.perf_counter()
284
- final_text = ""
285
 
286
- for partial in stream_generate(prompt):
287
- final_text = partial
288
- placeholder.markdown(final_text)
 
 
 
 
 
 
 
 
 
289
 
290
  gen_time = time.perf_counter() - start
291
- output_tokens = count_tokens(final_text)
292
  tps = output_tokens / gen_time if gen_time > 0 else 0.0
293
 
294
- live_msg = {
295
- "prompt": prompt,
296
- "search_time": search_time,
297
- "gen_time": gen_time,
298
- "input_tokens": input_tokens,
299
- "output_tokens": output_tokens,
300
- "tps": tps,
301
- }
302
-
303
- # Inline controls for the live message (thumbs disabled until it’s saved)
304
- render_inline_controls(
305
- msg=live_msg,
306
- feedback_key="fb_live",
307
- feedback_disabled=True,
308
- feedback_idx=None,
309
- )
 
 
 
310
 
311
- # Save assistant message so thumbs attach after rerun
312
  st.session_state.messages.append(
313
  {
314
  "role": "assistant",
315
- "content": final_text,
 
 
 
316
  "prompt": prompt,
317
  "search_time": search_time,
318
  "gen_time": gen_time,
319
  "input_tokens": input_tokens,
320
  "output_tokens": output_tokens,
321
  "tps": tps,
 
322
  "feedback": None,
323
  }
324
  )
 
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
 
 
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")
 
81
 
82
 
83
  # =========================
 
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
  # =========================
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
  # =========================
276
  # INPUT
 
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 (search) "
342
+ f"🧠 {gen_time:.2f}s (generation) "
343
+ f"⚡ {tps:.1f} tok/s "
344
+ f"🧾 {input_tokens} input tokens • {output_tokens} 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
  )