omarsol commited on
Commit
fa7c98c
·
1 Parent(s): 116fbd7

Improve retrieval streaming and tracing

Browse files
app/api.py CHANGED
@@ -344,6 +344,7 @@ class UIMessageStreamEncoder:
344
  self.active_reasoning_id = ""
345
  self.open_tool_call_ids: list[str] = []
346
  self.announced_tool_call_ids: set[str] = set()
 
347
  self.closed = False
348
 
349
  def close_reasoning_block(self) -> list[dict[str, Any]]:
@@ -463,7 +464,7 @@ class UIMessageStreamEncoder:
463
  )
464
  args = event.data.get("args")
465
  args_text = str(event.data.get("args_text", "")).strip()
466
- if isinstance(args, dict):
467
  parts.append(
468
  {
469
  "type": "tool-input-available",
@@ -472,6 +473,7 @@ class UIMessageStreamEncoder:
472
  "input": args,
473
  }
474
  )
 
475
  elif args_text:
476
  parts.append(
477
  {
@@ -481,6 +483,40 @@ class UIMessageStreamEncoder:
481
  "input": {"text": args_text},
482
  }
483
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
484
  return parts
485
 
486
  if event.type == "source_match":
@@ -507,16 +543,11 @@ class UIMessageStreamEncoder:
507
  parts.extend(self.close_reasoning_block())
508
  call_id = str(event.data.get("call_id", uuid4().hex))
509
  args = event.data.get("args")
510
- if (
511
- isinstance(args, dict) and args
512
- ) or call_id not in self.announced_tool_call_ids:
513
- # Providers that stream tool calls incrementally announce the
514
- # call before its args are parsed; refresh the input now that
515
- # the full args are known. A call id the stream never
516
- # announced (e.g. a ToolMessage with a missing id) must also
517
- # get an input part first: the AI SDK client throws on an
518
- # output for an unknown tool call and drops the rest of the
519
- # stream.
520
  parts.append(
521
  {
522
  "type": "tool-input-available",
@@ -526,6 +557,7 @@ class UIMessageStreamEncoder:
526
  }
527
  )
528
  self.announced_tool_call_ids.add(call_id)
 
529
  output = {
530
  "text": str(event.data.get("output_text", "")),
531
  "matches": [
 
344
  self.active_reasoning_id = ""
345
  self.open_tool_call_ids: list[str] = []
346
  self.announced_tool_call_ids: set[str] = set()
347
+ self.available_tool_call_ids: set[str] = set()
348
  self.closed = False
349
 
350
  def close_reasoning_block(self) -> list[dict[str, Any]]:
 
464
  )
465
  args = event.data.get("args")
466
  args_text = str(event.data.get("args_text", "")).strip()
467
+ if isinstance(args, dict) and args:
468
  parts.append(
469
  {
470
  "type": "tool-input-available",
 
473
  "input": args,
474
  }
475
  )
476
+ self.available_tool_call_ids.add(call_id)
477
  elif args_text:
478
  parts.append(
479
  {
 
483
  "input": {"text": args_text},
484
  }
485
  )
486
+ self.available_tool_call_ids.add(call_id)
487
+ return parts
488
+
489
+ if event.type == "tool_call_args_available":
490
+ call_id = str(event.data.get("call_id", uuid4().hex))
491
+ tool_name = str(event.data.get("tool_name", "tool"))
492
+ if call_id not in self.open_tool_call_ids:
493
+ self.open_tool_call_ids.append(call_id)
494
+ if call_id not in self.announced_tool_call_ids:
495
+ parts.append(
496
+ {
497
+ "type": "tool-input-start",
498
+ "toolCallId": call_id,
499
+ "toolName": tool_name,
500
+ }
501
+ )
502
+ self.announced_tool_call_ids.add(call_id)
503
+ args = event.data.get("args")
504
+ args_text = str(event.data.get("args_text", "")).strip()
505
+ if isinstance(args, dict):
506
+ input_data = args
507
+ elif args_text:
508
+ input_data = {"text": args_text}
509
+ else:
510
+ return parts
511
+ parts.append(
512
+ {
513
+ "type": "tool-input-available",
514
+ "toolCallId": call_id,
515
+ "toolName": tool_name,
516
+ "input": input_data,
517
+ }
518
+ )
519
+ self.available_tool_call_ids.add(call_id)
520
  return parts
521
 
522
  if event.type == "source_match":
 
543
  parts.extend(self.close_reasoning_block())
544
  call_id = str(event.data.get("call_id", uuid4().hex))
545
  args = event.data.get("args")
546
+ if call_id not in self.available_tool_call_ids:
547
+ # The completed-model update normally supplies full arguments
548
+ # before execution. Keep this completion-time fallback for
549
+ # providers that omit that update and for orphan ToolMessages:
550
+ # the AI SDK requires an input part before the output.
 
 
 
 
 
551
  parts.append(
552
  {
553
  "type": "tool-input-available",
 
557
  }
558
  )
559
  self.announced_tool_call_ids.add(call_id)
560
+ self.available_tool_call_ids.add(call_id)
561
  output = {
562
  "text": str(event.data.get("output_text", "")),
563
  "matches": [
app/chat_service.py CHANGED
@@ -688,6 +688,46 @@ def collect_retrieval_source_matches(payload: str) -> list[SourceMatch]:
688
  return matches
689
 
690
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
691
  def _record_evidence(
692
  target: dict[str, SourceMatch], matches: list[SourceMatch]
693
  ) -> None:
@@ -1358,6 +1398,15 @@ class StableToolOutputCapMiddleware(AgentMiddleware):
1358
  "sha256": hashlib.sha256(raw).hexdigest(),
1359
  "max_bytes": self.max_bytes,
1360
  }
 
 
 
 
 
 
 
 
 
1361
  turn = _turn_id_for(request)
1362
  record_turn_signal(turn, "tool_outputs_capped", 1)
1363
  record_turn_signal(turn, "tool_output_original_bytes", len(raw))
@@ -2894,7 +2943,11 @@ async def stream_chat(request: ChatRequest) -> AsyncIterator[ChatEvent]:
2894
  tool_name = str(getattr(message, "name", "tool"))
2895
  call_matches: list[SourceMatch] = []
2896
  if tool_name == "retrieve_tutor_context":
2897
- call_matches = collect_retrieval_source_matches(payload)
 
 
 
 
2898
  _record_evidence(retrieval_evidence, call_matches)
2899
 
2900
  tool_call = tool_calls_by_id.get(
@@ -3009,10 +3062,44 @@ async def stream_chat(request: ChatRequest) -> AsyncIterator[ChatEvent]:
3009
  if getattr(message, "tool_calls", None):
3010
  # The completed message has fully parsed args; streamed
3011
  # fragments may have announced the call with empty args.
 
 
 
3012
  for tool_call in message.tool_calls:
3013
  call_id = str(tool_call.get("id") or "")
3014
- if call_id:
3015
- tool_calls_by_id[call_id] = tool_call
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3016
  continue
3017
  completed_answer = message_content_to_text(message.content)
3018
  finally:
 
688
  return matches
689
 
690
 
691
+ def _source_match_record(match: SourceMatch) -> dict[str, Any]:
692
+ """Store lightweight source metadata alongside a persistently capped tool."""
693
+ return {
694
+ "doc_id": match.doc_id,
695
+ "title": match.title,
696
+ "url": match.url,
697
+ "source_key": match.source_key,
698
+ "source_label": match.source_label,
699
+ "score": match.score,
700
+ "group": match.group,
701
+ "path": match.path,
702
+ }
703
+
704
+
705
+ def _source_matches_from_records(records: Any) -> list[SourceMatch]:
706
+ """Restore source metadata retained outside a truncated retrieval payload."""
707
+ if not isinstance(records, list):
708
+ return []
709
+ matches: list[SourceMatch] = []
710
+ for record in records:
711
+ if not isinstance(record, dict):
712
+ continue
713
+ try:
714
+ matches.append(
715
+ SourceMatch(
716
+ doc_id=str(record.get("doc_id", "")),
717
+ title=str(record.get("title", "")),
718
+ url=str(record.get("url", "")),
719
+ source_key=str(record.get("source_key", "")),
720
+ source_label=str(record.get("source_label", "")),
721
+ score=float(record.get("score", 0.0)),
722
+ group=str(record.get("group", "")),
723
+ path=str(record.get("path", "")),
724
+ )
725
+ )
726
+ except (TypeError, ValueError):
727
+ continue
728
+ return matches
729
+
730
+
731
  def _record_evidence(
732
  target: dict[str, SourceMatch], matches: list[SourceMatch]
733
  ) -> None:
 
1398
  "sha256": hashlib.sha256(raw).hexdigest(),
1399
  "max_bytes": self.max_bytes,
1400
  }
1401
+ if getattr(result, "name", "") == "retrieve_tutor_context":
1402
+ # Capping inserts a head/tail marker and intentionally makes the
1403
+ # JSON content unparsable. Preserve only lightweight source
1404
+ # metadata so citations and the UI's chunk count stay accurate
1405
+ # without keeping the discarded chunk bodies in the checkpoint.
1406
+ retrieval_matches = collect_retrieval_source_matches(text)
1407
+ metadata["retrieval_matches"] = [
1408
+ _source_match_record(match) for match in retrieval_matches
1409
+ ]
1410
  turn = _turn_id_for(request)
1411
  record_turn_signal(turn, "tool_outputs_capped", 1)
1412
  record_turn_signal(turn, "tool_output_original_bytes", len(raw))
 
2943
  tool_name = str(getattr(message, "name", "tool"))
2944
  call_matches: list[SourceMatch] = []
2945
  if tool_name == "retrieve_tutor_context":
2946
+ call_matches = _source_matches_from_records(
2947
+ cap_metadata.get("retrieval_matches")
2948
+ )
2949
+ if not call_matches:
2950
+ call_matches = collect_retrieval_source_matches(payload)
2951
  _record_evidence(retrieval_evidence, call_matches)
2952
 
2953
  tool_call = tool_calls_by_id.get(
 
3062
  if getattr(message, "tool_calls", None):
3063
  # The completed message has fully parsed args; streamed
3064
  # fragments may have announced the call with empty args.
3065
+ # Publish the completed arguments now, before the tools
3066
+ # node returns, so a long-running search shows its query
3067
+ # while it is running rather than only with its result.
3068
  for tool_call in message.tool_calls:
3069
  call_id = str(tool_call.get("id") or "")
3070
+ if not call_id:
3071
+ continue
3072
+ previous = tool_calls_by_id.get(call_id)
3073
+ tool_calls_by_id[call_id] = tool_call
3074
+ if previous is None:
3075
+ if first_token_at is None:
3076
+ first_token_at = time.monotonic()
3077
+ yield ChatEvent(
3078
+ "tool_call_started",
3079
+ {
3080
+ "message_id": message_id,
3081
+ "call_id": call_id,
3082
+ "tool_name": str(tool_call.get("name", "tool")),
3083
+ "args": tool_call.get("args"),
3084
+ "args_text": format_tool_args(
3085
+ tool_call.get("args")
3086
+ ),
3087
+ },
3088
+ )
3089
+ continue
3090
+ args = tool_call.get("args")
3091
+ if args == previous.get("args"):
3092
+ continue
3093
+ yield ChatEvent(
3094
+ "tool_call_args_available",
3095
+ {
3096
+ "message_id": message_id,
3097
+ "call_id": call_id,
3098
+ "tool_name": str(tool_call.get("name", "tool")),
3099
+ "args": args,
3100
+ "args_text": format_tool_args(args),
3101
+ },
3102
+ )
3103
  continue
3104
  completed_answer = message_content_to_text(message.content)
3105
  finally:
app/chroma_rag.py CHANGED
@@ -11,14 +11,17 @@ import re
11
  import threading
12
  import time
13
  from collections import Counter, deque
 
14
  from dataclasses import asdict, dataclass
15
  from pathlib import Path
16
- from typing import Any, Iterable
17
  from urllib.parse import unquote, urlparse
18
 
19
  import chromadb
20
  import cohere
21
  import tiktoken
 
 
22
  from tqdm.auto import tqdm
23
 
24
  logger = logging.getLogger(__name__)
@@ -164,6 +167,78 @@ class SearchResult:
164
  retrieval_method: str = ""
165
 
166
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
167
  @dataclass(slots=True)
168
  class BM25Index:
169
  records: list[ChunkRecord]
@@ -842,7 +917,14 @@ def _wait_for_cohere_retry(
842
  else:
843
  delay = min(window_seconds, max(15.0, 2.0**attempt))
844
 
845
- time.sleep(delay + random.uniform(0.5, 2.0))
 
 
 
 
 
 
 
846
 
847
 
848
  def _cohere_embeddings_list(response: Any) -> list[list[float]]:
@@ -1253,8 +1335,50 @@ class LocalChromaRetriever:
1253
  self._document_dict: dict[str, dict[str, Any]] = pickle.load(handle)
1254
 
1255
  self._bm25_index = load_bm25_index(self._bm25_index_path)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1256
  self._cohere = cohere.ClientV2(api_key=cohere_api_key)
1257
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1258
  def search(
1259
  self,
1260
  query: str,
@@ -1262,67 +1386,200 @@ class LocalChromaRetriever:
1262
  allowed_sources: list[str] | None = None,
1263
  token_budget: int | None = None,
1264
  ) -> list[SearchResult]:
1265
- dense_hits = self._dense_search(query, allowed_sources=allowed_sources)
1266
- bm25_hits = self._bm25_search(query, allowed_sources=allowed_sources)
1267
- fused_hits = reciprocal_rank_fusion(
1268
- [hits for hits in (dense_hits, bm25_hits) if hits],
1269
- rrf_k=self._rrf_k,
1270
- top_k=self._fusion_top_k,
1271
  )
1272
- if not fused_hits:
1273
- return []
 
 
 
 
 
 
 
1274
 
1275
- reranked = rerank_results(
1276
- self._cohere,
1277
- query,
1278
- fused_hits,
1279
- model=self._rerank_model,
1280
- top_n=self._rerank_top_k,
1281
- )
1282
- return self._apply_token_budget(reranked, token_budget)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1283
 
1284
  def _dense_search(
1285
  self,
1286
  query: str,
1287
  *,
1288
  allowed_sources: list[str] | None = None,
 
1289
  ) -> list[SearchResult]:
1290
- query_embedding = embed_texts(
1291
- self._cohere,
1292
- [query],
1293
- input_type="search_query",
1294
- model=self._embed_model,
1295
- )[0]
 
 
 
 
 
 
1296
 
1297
  where = build_where_filter(allowed_sources)
1298
- raw_results = self._collection.query(
1299
- query_embeddings=[query_embedding],
1300
- n_results=self._dense_top_k,
1301
- where=where,
1302
- include=["documents", "metadatas", "distances"],
1303
- )
1304
-
1305
- chunk_ids = _flatten_query_results(raw_results.get("ids"))
1306
- documents = _flatten_query_results(raw_results.get("documents"))
1307
- metadatas = _flatten_query_results(raw_results.get("metadatas"))
1308
- distances = _flatten_query_results(raw_results.get("distances"))
 
 
 
 
1309
 
1310
- dense_hits: list[SearchResult] = []
1311
- for chunk_id, chunk_text, metadata, distance in zip(
1312
- chunk_ids, documents, metadatas, distances, strict=False
 
1313
  ):
1314
- if metadata is None:
1315
- continue
 
 
 
 
 
 
 
 
 
1316
 
1317
- dense_hits.append(
1318
- self._search_result_from_metadata(
1319
- chunk_id=str(chunk_id),
1320
- score=_distance_to_score(distance),
1321
- chunk_text=str(chunk_text),
1322
- metadata=dict(metadata),
1323
- retrieval_method="dense",
 
1324
  )
1325
- )
1326
  return dense_hits
1327
 
1328
  def _bm25_search(
 
11
  import threading
12
  import time
13
  from collections import Counter, deque
14
+ from contextlib import contextmanager
15
  from dataclasses import asdict, dataclass
16
  from pathlib import Path
17
+ from typing import Any, Iterable, Iterator
18
  from urllib.parse import unquote, urlparse
19
 
20
  import chromadb
21
  import cohere
22
  import tiktoken
23
+ from langsmith import trace, traceable
24
+ from langsmith.run_helpers import get_current_run_tree
25
  from tqdm.auto import tqdm
26
 
27
  logger = logging.getLogger(__name__)
 
167
  retrieval_method: str = ""
168
 
169
 
170
+ RETRIEVAL_STAGE_NAMES = (
171
+ "embed_ms",
172
+ "chroma_ms",
173
+ "dense_hydration_ms",
174
+ "bm25_ms",
175
+ "fusion_ms",
176
+ "rerank_ms",
177
+ "token_budget_ms",
178
+ )
179
+
180
+ RETRIEVAL_STAGE_TRACE_NAMES = {
181
+ "embed_ms": "Cohere Embed",
182
+ "chroma_ms": "Chroma Vector Search",
183
+ "dense_hydration_ms": "Dense Result Hydration",
184
+ "bm25_ms": "BM25 Search",
185
+ "fusion_ms": "RRF Fusion",
186
+ "rerank_ms": "Cohere Rerank",
187
+ "token_budget_ms": "Token Budget",
188
+ }
189
+
190
+
191
+ @contextmanager
192
+ def _measure_retrieval_stage(
193
+ timings: dict[str, float],
194
+ stage: str,
195
+ *,
196
+ trace_inputs: dict[str, Any] | None = None,
197
+ ) -> Iterator[None]:
198
+ started_at = time.perf_counter()
199
+ with trace(
200
+ RETRIEVAL_STAGE_TRACE_NAMES[stage],
201
+ run_type="chain",
202
+ inputs=trace_inputs or {},
203
+ metadata={"retrieval_stage": stage},
204
+ ) as stage_run:
205
+ try:
206
+ yield
207
+ except BaseException:
208
+ elapsed_ms = (time.perf_counter() - started_at) * 1000
209
+ timings[stage] = elapsed_ms
210
+ stage_run.add_metadata({"duration_ms": round(elapsed_ms, 2)})
211
+ raise
212
+ else:
213
+ elapsed_ms = (time.perf_counter() - started_at) * 1000
214
+ timings[stage] = elapsed_ms
215
+ stage_run.end(metadata={"duration_ms": round(elapsed_ms, 2)})
216
+
217
+
218
+ def _trace_retrieval_inputs(inputs: dict[str, Any]) -> dict[str, Any]:
219
+ allowed_sources = inputs.get("allowed_sources")
220
+ return {
221
+ "query": str(inputs.get("query", "")),
222
+ "requested_source_count": len(set(allowed_sources or [])),
223
+ "token_budget": inputs.get("token_budget"),
224
+ }
225
+
226
+
227
+ def _trace_retrieval_outputs(results: Any) -> dict[str, Any]:
228
+ if not isinstance(results, list):
229
+ return {"result_type": type(results).__name__}
230
+ return {
231
+ "result_count": len(results),
232
+ "sources": sorted(
233
+ {
234
+ result.source
235
+ for result in results
236
+ if isinstance(result, SearchResult) and result.source
237
+ }
238
+ ),
239
+ }
240
+
241
+
242
  @dataclass(slots=True)
243
  class BM25Index:
244
  records: list[ChunkRecord]
 
917
  else:
918
  delay = min(window_seconds, max(15.0, 2.0**attempt))
919
 
920
+ sleep_seconds = delay + random.uniform(0.5, 2.0)
921
+ logger.warning(
922
+ "cohere_embed_retry attempt=%d delay_seconds=%.2f status_code=%s",
923
+ attempt,
924
+ sleep_seconds,
925
+ getattr(exc, "status_code", 429),
926
+ )
927
+ time.sleep(sleep_seconds)
928
 
929
 
930
  def _cohere_embeddings_list(response: Any) -> list[list[float]]:
 
1335
  self._document_dict: dict[str, dict[str, Any]] = pickle.load(handle)
1336
 
1337
  self._bm25_index = load_bm25_index(self._bm25_index_path)
1338
+ document_sources = {
1339
+ str(document.get("source", "")).strip()
1340
+ for document in self._document_dict.values()
1341
+ if isinstance(document, dict) and document.get("source")
1342
+ }
1343
+ bm25_sources = (
1344
+ {
1345
+ str(record.metadata.get("source", "")).strip()
1346
+ for record in self._bm25_index.records
1347
+ if record.metadata.get("source")
1348
+ }
1349
+ if self._bm25_index is not None
1350
+ else set()
1351
+ )
1352
+ # These artifacts are loaded once by the process-cached retriever. Use
1353
+ # their actual source set instead of a configured UI default so a
1354
+ # complete source selection can safely skip Chroma's redundant,
1355
+ # expensive all-record metadata filter.
1356
+ self._indexed_sources = frozenset(document_sources | bm25_sources)
1357
  self._cohere = cohere.ClientV2(api_key=cohere_api_key)
1358
 
1359
+ def _effective_allowed_sources(
1360
+ self, allowed_sources: list[str] | None
1361
+ ) -> tuple[list[str] | None, str]:
1362
+ if allowed_sources is None:
1363
+ return None, "unfiltered"
1364
+
1365
+ requested_sources = frozenset(allowed_sources)
1366
+ if not requested_sources:
1367
+ # The normal chat path disables local tools for an explicit empty
1368
+ # source selection. Preserve the existing helper semantics here.
1369
+ return allowed_sources, "empty_selection"
1370
+ if self._indexed_sources and self._indexed_sources.issubset(requested_sources):
1371
+ return None, "all_sources_omitted"
1372
+ if len(requested_sources) == 1:
1373
+ return allowed_sources, "single_source"
1374
+ return allowed_sources, "source_subset"
1375
+
1376
+ @traceable(
1377
+ name="Hybrid Retrieval Pipeline",
1378
+ run_type="retriever",
1379
+ process_inputs=_trace_retrieval_inputs,
1380
+ process_outputs=_trace_retrieval_outputs,
1381
+ )
1382
  def search(
1383
  self,
1384
  query: str,
 
1386
  allowed_sources: list[str] | None = None,
1387
  token_budget: int | None = None,
1388
  ) -> list[SearchResult]:
1389
+ started_at = time.perf_counter()
1390
+ timings = {stage: 0.0 for stage in RETRIEVAL_STAGE_NAMES}
1391
+ effective_sources, filter_mode = self._effective_allowed_sources(
1392
+ allowed_sources
 
 
1393
  )
1394
+ requested_source_count = len(set(allowed_sources or []))
1395
+ counts = {
1396
+ "dense_hit_count": 0,
1397
+ "bm25_hit_count": 0,
1398
+ "fused_hit_count": 0,
1399
+ "reranked_hit_count": 0,
1400
+ "result_count": 0,
1401
+ }
1402
+ status = "success"
1403
 
1404
+ try:
1405
+ dense_hits = self._dense_search(
1406
+ query,
1407
+ allowed_sources=effective_sources,
1408
+ timings=timings,
1409
+ )
1410
+ counts["dense_hit_count"] = len(dense_hits)
1411
+
1412
+ with _measure_retrieval_stage(
1413
+ timings,
1414
+ "bm25_ms",
1415
+ trace_inputs={
1416
+ "top_k": self._bm25_top_k,
1417
+ "filter_applied": effective_sources is not None,
1418
+ "source_count": len(set(effective_sources or [])),
1419
+ },
1420
+ ):
1421
+ bm25_hits = self._bm25_search(query, allowed_sources=effective_sources)
1422
+ counts["bm25_hit_count"] = len(bm25_hits)
1423
+
1424
+ with _measure_retrieval_stage(
1425
+ timings,
1426
+ "fusion_ms",
1427
+ trace_inputs={
1428
+ "dense_hit_count": len(dense_hits),
1429
+ "bm25_hit_count": len(bm25_hits),
1430
+ "top_k": self._fusion_top_k,
1431
+ "rrf_k": self._rrf_k,
1432
+ },
1433
+ ):
1434
+ fused_hits = reciprocal_rank_fusion(
1435
+ [hits for hits in (dense_hits, bm25_hits) if hits],
1436
+ rrf_k=self._rrf_k,
1437
+ top_k=self._fusion_top_k,
1438
+ )
1439
+ counts["fused_hit_count"] = len(fused_hits)
1440
+ if not fused_hits:
1441
+ return []
1442
+
1443
+ with _measure_retrieval_stage(
1444
+ timings,
1445
+ "rerank_ms",
1446
+ trace_inputs={
1447
+ "model": self._rerank_model,
1448
+ "candidate_count": len(fused_hits),
1449
+ "top_n": self._rerank_top_k,
1450
+ },
1451
+ ):
1452
+ reranked = rerank_results(
1453
+ self._cohere,
1454
+ query,
1455
+ fused_hits,
1456
+ model=self._rerank_model,
1457
+ top_n=self._rerank_top_k,
1458
+ )
1459
+ counts["reranked_hit_count"] = len(reranked)
1460
+
1461
+ with _measure_retrieval_stage(
1462
+ timings,
1463
+ "token_budget_ms",
1464
+ trace_inputs={
1465
+ "candidate_count": len(reranked),
1466
+ "token_budget": (
1467
+ self._token_budget if token_budget is None else token_budget
1468
+ ),
1469
+ },
1470
+ ):
1471
+ results = self._apply_token_budget(reranked, token_budget)
1472
+ counts["result_count"] = len(results)
1473
+ return results
1474
+ except Exception:
1475
+ status = "error"
1476
+ raise
1477
+ finally:
1478
+ total_ms = (time.perf_counter() - started_at) * 1000
1479
+ timing_metadata: dict[str, Any] = {
1480
+ "status": status,
1481
+ "filter_mode": filter_mode,
1482
+ "requested_source_count": requested_source_count,
1483
+ "indexed_source_count": len(self._indexed_sources),
1484
+ "total_ms": round(total_ms, 2),
1485
+ "stage_ms": {
1486
+ stage: round(timings[stage], 2) for stage in RETRIEVAL_STAGE_NAMES
1487
+ },
1488
+ **counts,
1489
+ }
1490
+ current_run = get_current_run_tree()
1491
+ if current_run is not None:
1492
+ current_run.add_metadata({"retrieval_timing": timing_metadata})
1493
+
1494
+ logger.info(
1495
+ "retrieval_timing status=%s filter_mode=%s "
1496
+ "requested_sources=%d indexed_sources=%d total_ms=%.2f "
1497
+ "embed_ms=%.2f chroma_ms=%.2f dense_hydration_ms=%.2f "
1498
+ "bm25_ms=%.2f fusion_ms=%.2f rerank_ms=%.2f "
1499
+ "token_budget_ms=%.2f dense_hits=%d bm25_hits=%d "
1500
+ "fused_hits=%d reranked_hits=%d results=%d",
1501
+ status,
1502
+ filter_mode,
1503
+ requested_source_count,
1504
+ len(self._indexed_sources),
1505
+ total_ms,
1506
+ timings["embed_ms"],
1507
+ timings["chroma_ms"],
1508
+ timings["dense_hydration_ms"],
1509
+ timings["bm25_ms"],
1510
+ timings["fusion_ms"],
1511
+ timings["rerank_ms"],
1512
+ timings["token_budget_ms"],
1513
+ counts["dense_hit_count"],
1514
+ counts["bm25_hit_count"],
1515
+ counts["fused_hit_count"],
1516
+ counts["reranked_hit_count"],
1517
+ counts["result_count"],
1518
+ )
1519
 
1520
  def _dense_search(
1521
  self,
1522
  query: str,
1523
  *,
1524
  allowed_sources: list[str] | None = None,
1525
+ timings: dict[str, float] | None = None,
1526
  ) -> list[SearchResult]:
1527
+ stage_timings = timings if timings is not None else {}
1528
+ with _measure_retrieval_stage(
1529
+ stage_timings,
1530
+ "embed_ms",
1531
+ trace_inputs={"model": self._embed_model, "input_count": 1},
1532
+ ):
1533
+ query_embedding = embed_texts(
1534
+ self._cohere,
1535
+ [query],
1536
+ input_type="search_query",
1537
+ model=self._embed_model,
1538
+ )[0]
1539
 
1540
  where = build_where_filter(allowed_sources)
1541
+ with _measure_retrieval_stage(
1542
+ stage_timings,
1543
+ "chroma_ms",
1544
+ trace_inputs={
1545
+ "collection": self._collection_name,
1546
+ "top_k": self._dense_top_k,
1547
+ "where": where,
1548
+ },
1549
+ ):
1550
+ raw_results = self._collection.query(
1551
+ query_embeddings=[query_embedding],
1552
+ n_results=self._dense_top_k,
1553
+ where=where,
1554
+ include=["documents", "metadatas", "distances"],
1555
+ )
1556
 
1557
+ with _measure_retrieval_stage(
1558
+ stage_timings,
1559
+ "dense_hydration_ms",
1560
+ trace_inputs={"requested_top_k": self._dense_top_k},
1561
  ):
1562
+ chunk_ids = _flatten_query_results(raw_results.get("ids"))
1563
+ documents = _flatten_query_results(raw_results.get("documents"))
1564
+ metadatas = _flatten_query_results(raw_results.get("metadatas"))
1565
+ distances = _flatten_query_results(raw_results.get("distances"))
1566
+
1567
+ dense_hits: list[SearchResult] = []
1568
+ for chunk_id, chunk_text, metadata, distance in zip(
1569
+ chunk_ids, documents, metadatas, distances, strict=False
1570
+ ):
1571
+ if metadata is None:
1572
+ continue
1573
 
1574
+ dense_hits.append(
1575
+ self._search_result_from_metadata(
1576
+ chunk_id=str(chunk_id),
1577
+ score=_distance_to_score(distance),
1578
+ chunk_text=str(chunk_text),
1579
+ metadata=dict(metadata),
1580
+ retrieval_method="dense",
1581
+ )
1582
  )
 
1583
  return dense_hits
1584
 
1585
  def _bm25_search(
frontend/components/chat-message.test.tsx CHANGED
@@ -54,6 +54,7 @@ describe("ChatMessage activity rendering", () => {
54
  expect(items[0]?.textContent).toContain("Continuation of the first thought");
55
  expect(items[1]?.textContent).toContain("Hybrid search");
56
  expect(items[1]?.textContent).toContain("agent memory");
 
57
  expect(items[2]?.textContent).toContain("Second thought after retrieval");
58
  expect(screen.getByText("Final answer")).toBeTruthy();
59
  });
 
54
  expect(items[0]?.textContent).toContain("Continuation of the first thought");
55
  expect(items[1]?.textContent).toContain("Hybrid search");
56
  expect(items[1]?.textContent).toContain("agent memory");
57
+ expect(items[1]?.textContent).toContain("3 chunks");
58
  expect(items[2]?.textContent).toContain("Second thought after retrieval");
59
  expect(screen.getByText("Final answer")).toBeTruthy();
60
  });
frontend/components/chat-message.tsx CHANGED
@@ -463,6 +463,7 @@ function ToolRow({ part }: { part: TutorMessagePart }) {
463
  ? outputObject.matches.length
464
  : 0;
465
  const resultSummary = formatToolResultSummary({
 
466
  outputText,
467
  matchCount,
468
  state: part.state,
@@ -600,11 +601,13 @@ function formatToolStateBadge(
600
  }
601
 
602
  function formatToolResultSummary({
 
603
  outputText,
604
  matchCount,
605
  state,
606
  errorText,
607
  }: {
 
608
  outputText: string;
609
  matchCount: number;
610
  state?: string;
@@ -614,7 +617,10 @@ function formatToolResultSummary({
614
  return "";
615
  }
616
  if (matchCount > 0) {
617
- return `${matchCount} match${matchCount === 1 ? "" : "es"}`;
 
 
 
618
  }
619
  if (outputText) {
620
  const lineCount = outputText.split("\n").length;
 
463
  ? outputObject.matches.length
464
  : 0;
465
  const resultSummary = formatToolResultSummary({
466
+ toolType: part.type,
467
  outputText,
468
  matchCount,
469
  state: part.state,
 
601
  }
602
 
603
  function formatToolResultSummary({
604
+ toolType,
605
  outputText,
606
  matchCount,
607
  state,
608
  errorText,
609
  }: {
610
+ toolType: string;
611
  outputText: string;
612
  matchCount: number;
613
  state?: string;
 
617
  return "";
618
  }
619
  if (matchCount > 0) {
620
+ const isRetrieval =
621
+ toolType.replace(/^tool-/, "") === "retrieve_tutor_context";
622
+ const noun = isRetrieval ? "chunk" : "match";
623
+ return `${matchCount} ${noun}${matchCount === 1 ? "" : "s"}`;
624
  }
625
  if (outputText) {
626
  const lineCount = outputText.split("\n").length;
tests/conftest.py ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import os
4
+
5
+
6
+ # app.config loads the repository's .env during test collection. Without an
7
+ # explicit override, ordinary unit tests upload synthetic LangGraph and
8
+ # retriever runs into the production LangSmith project. Keep normal pytest
9
+ # hermetic; the opt-in live E2E mode deliberately retains tracing.
10
+ if os.getenv("RUN_LIVE_API_E2E") != "1":
11
+ os.environ["LANGSMITH_TRACING"] = "false"
12
+ os.environ["LANGSMITH_TRACING_V2"] = "false"
13
+ os.environ["LANGCHAIN_TRACING_V2"] = "false"
tests/manual_e2e_langsmith.md CHANGED
@@ -176,6 +176,10 @@ Expected result:
176
  - The stream includes `tool-input-*`, `tool-output-available`, `text-delta`,
177
  `source-url`, `source-document`, `data-source`, and `finish` parts.
178
  - Tool calls include `retrieve_tutor_context` and `run_kb_command`.
 
 
 
 
179
  - The answer has inline citations, not only a final sources list.
180
  - `data-source` parts include the source cards the frontend will render.
181
 
@@ -266,6 +270,45 @@ jq -r '
266
  ' "/tmp/ai_tutor_trace_${TRACE_ID}.json"
267
  ```
268
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
269
  LLM calls:
270
 
271
  ```bash
@@ -295,9 +338,12 @@ For latency debugging, compare these numbers:
295
  - Root trace duration.
296
  - Count of `run_kb_command` calls.
297
  - Total shell/tool duration.
 
298
  - Count of LLM/chat model calls.
299
  - Longest LLM call duration.
300
  - Whether the trace ended in `success`, `error`, or cancellation.
301
 
302
- If shell/tool duration is tiny but total trace duration is large, the bottleneck
303
- is model inference or repeated model turns, not filesystem access.
 
 
 
176
  - The stream includes `tool-input-*`, `tool-output-available`, `text-delta`,
177
  `source-url`, `source-document`, `data-source`, and `finish` parts.
178
  - Tool calls include `retrieve_tutor_context` and `run_kb_command`.
179
+ - Every `retrieve_tutor_context` tool run has a nested
180
+ `Hybrid Retrieval Pipeline` span. Expanding it in the waterfall shows
181
+ separate Cohere embed, Chroma, dense hydration, BM25, RRF, Cohere rerank,
182
+ and token-budget runs.
183
  - The answer has inline citations, not only a final sources list.
184
  - `data-source` parts include the source cards the frontend will render.
185
 
 
270
  ' "/tmp/ai_tutor_trace_${TRACE_ID}.json"
271
  ```
272
 
273
+ Hybrid retrieval latency breakdown:
274
+
275
+ ```bash
276
+ jq -r '
277
+ .runs[]
278
+ | select(.name == "Hybrid Retrieval Pipeline")
279
+ | (.custom_metadata.retrieval_timing
280
+ // .extra.metadata.retrieval_timing
281
+ // {}) as $timing
282
+ | [
283
+ (.inputs.query // ""),
284
+ ($timing.status // ""),
285
+ ($timing.filter_mode // ""),
286
+ ($timing.total_ms // ""),
287
+ ($timing.stage_ms.embed_ms // ""),
288
+ ($timing.stage_ms.chroma_ms // ""),
289
+ ($timing.stage_ms.dense_hydration_ms // ""),
290
+ ($timing.stage_ms.bm25_ms // ""),
291
+ ($timing.stage_ms.fusion_ms // ""),
292
+ ($timing.stage_ms.rerank_ms // ""),
293
+ ($timing.stage_ms.token_budget_ms // "")
294
+ ]
295
+ | @tsv
296
+ ' "/tmp/ai_tutor_trace_${TRACE_ID}.json"
297
+ ```
298
+
299
+ The columns are query, status, filter mode, total, embed, Chroma, dense-result
300
+ hydration, BM25, RRF, rerank, and token-budget latency, all in milliseconds.
301
+ The backend logs the same values as one `retrieval_timing` line without logging
302
+ the raw student query. A complete source selection should report
303
+ `filter_mode=all_sources_omitted`; real source subsets should report
304
+ `single_source` or `source_subset`.
305
+
306
+ In the LangSmith waterfall, expand `retrieve_tutor_context`, then
307
+ `Hybrid Retrieval Pipeline`, to see the same stages as individually timed child
308
+ runs. Their names are `Cohere Embed`, `Chroma Vector Search`,
309
+ `Dense Result Hydration`, `BM25 Search`, `RRF Fusion`, `Cohere Rerank`, and
310
+ `Token Budget`.
311
+
312
  LLM calls:
313
 
314
  ```bash
 
338
  - Root trace duration.
339
  - Count of `run_kb_command` calls.
340
  - Total shell/tool duration.
341
+ - The `Hybrid Retrieval Pipeline` stage breakdown for every retrieval call.
342
  - Count of LLM/chat model calls.
343
  - Longest LLM call duration.
344
  - Whether the trace ended in `success`, `error`, or cancellation.
345
 
346
+ If `Hybrid Retrieval Pipeline` is fast but total trace duration is large, the
347
+ bottleneck is model inference or repeated model turns. If retrieval is slow,
348
+ its child runs and metadata identify whether the time was spent in Cohere
349
+ embedding, Chroma, BM25/RRF, Cohere reranking, or token budgeting.
tests/test_api.py CHANGED
@@ -612,6 +612,60 @@ class ApiTestCase(unittest.TestCase):
612
  self.assertEqual(matches[0]["score"], 0.9)
613
  self.assertEqual(matches[0]["path"], "raw/docs/peft/lora.md")
614
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
615
  def test_chat_rejects_oversized_query(self) -> None:
616
  from app.api import MAX_QUERY_CHARS
617
 
 
612
  self.assertEqual(matches[0]["score"], 0.9)
613
  self.assertEqual(matches[0]["path"], "raw/docs/peft/lora.md")
614
 
615
+ def test_completed_tool_args_refresh_while_the_tool_is_running(self) -> None:
616
+ encoder = UIMessageStreamEncoder()
617
+ encoder.encode(ChatEvent("message_started", {"message_id": "m1"}))
618
+
619
+ started = encoder.encode(
620
+ ChatEvent(
621
+ "tool_call_started",
622
+ {
623
+ "message_id": "m1",
624
+ "call_id": "call_1",
625
+ "tool_name": "retrieve_tutor_context",
626
+ "args": {},
627
+ },
628
+ )
629
+ )
630
+ self.assertEqual([part["type"] for part in started], ["tool-input-start"])
631
+
632
+ refreshed = encoder.encode(
633
+ ChatEvent(
634
+ "tool_call_args_available",
635
+ {
636
+ "message_id": "m1",
637
+ "call_id": "call_1",
638
+ "tool_name": "retrieve_tutor_context",
639
+ "args": {"query": "secure API key storage"},
640
+ },
641
+ )
642
+ )
643
+ self.assertEqual(
644
+ [part["type"] for part in refreshed],
645
+ ["tool-input-available"],
646
+ )
647
+ self.assertEqual(
648
+ refreshed[0]["input"],
649
+ {"query": "secure API key storage"},
650
+ )
651
+
652
+ completed = encoder.encode(
653
+ ChatEvent(
654
+ "tool_call_completed",
655
+ {
656
+ "message_id": "m1",
657
+ "call_id": "call_1",
658
+ "tool_name": "retrieve_tutor_context",
659
+ "args": {"query": "secure API key storage"},
660
+ "output_text": "payload",
661
+ },
662
+ )
663
+ )
664
+ self.assertNotIn(
665
+ "tool-input-available",
666
+ [part["type"] for part in completed],
667
+ )
668
+
669
  def test_chat_rejects_oversized_query(self) -> None:
670
  from app.api import MAX_QUERY_CHARS
671
 
tests/test_chat_service.py CHANGED
@@ -117,6 +117,90 @@ class FakeStreamingAgent(FakeAgent):
117
  }
118
 
119
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
120
  class FakeAnswerAgent(FakeAgent):
121
  """Agent that streams nothing and answers with one final AI message."""
122
 
@@ -871,6 +955,51 @@ class ChatServiceTestCase(unittest.TestCase):
871
  )
872
  self.assertTrue(build_agent_mock.call_args.kwargs["include_thoughts"])
873
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
874
  def test_stream_chat_resolves_shell_citation_after_final_answer(self) -> None:
875
  agent = FakeStreamingAgent([])
876
  self.addCleanup(_drop_thread_record, "thread_rg")
 
117
  }
118
 
119
 
120
+ class FakeIncrementalRetrievalAgent(FakeAgent):
121
+ """DeepSeek-like stream: tool id/name first, complete args at model end."""
122
+
123
+ async def astream(self, *args, **kwargs):
124
+ yield {
125
+ "type": "messages",
126
+ "data": (
127
+ AIMessageChunk(
128
+ content="",
129
+ tool_calls=[
130
+ {
131
+ "id": "call_retrieval",
132
+ "name": "retrieve_tutor_context",
133
+ "args": {},
134
+ }
135
+ ],
136
+ ),
137
+ {"langgraph_node": "model"},
138
+ ),
139
+ }
140
+ yield {
141
+ "type": "updates",
142
+ "data": {
143
+ "model": {
144
+ "messages": [
145
+ AIMessage(
146
+ content="",
147
+ tool_calls=[
148
+ {
149
+ "id": "call_retrieval",
150
+ "name": "retrieve_tutor_context",
151
+ "args": {"query": "secure API key storage"},
152
+ }
153
+ ],
154
+ )
155
+ ]
156
+ }
157
+ },
158
+ }
159
+ yield {
160
+ "type": "updates",
161
+ "data": {
162
+ "tools": {
163
+ "messages": [
164
+ ToolMessage(
165
+ content=(
166
+ '{"query":"secure API key storage","matches":[...\n\n'
167
+ "[... tool output truncated ...]\n\n...]}"
168
+ ),
169
+ name="retrieve_tutor_context",
170
+ tool_call_id="call_retrieval",
171
+ additional_kwargs={
172
+ "stable_tool_cap": {
173
+ "retrieval_matches": [
174
+ {
175
+ "doc_id": "doc-keys",
176
+ "title": "Manage API keys",
177
+ "url": "https://example.com/keys",
178
+ "source_key": "langchain",
179
+ "source_label": "LangChain Docs",
180
+ "score": 0.9,
181
+ "group": "docs",
182
+ "path": "",
183
+ }
184
+ ]
185
+ }
186
+ },
187
+ )
188
+ ]
189
+ }
190
+ },
191
+ }
192
+ yield {
193
+ "type": "updates",
194
+ "data": {
195
+ "model": {
196
+ "messages": [
197
+ AIMessage(content="Store secrets outside source control.")
198
+ ]
199
+ }
200
+ },
201
+ }
202
+
203
+
204
  class FakeAnswerAgent(FakeAgent):
205
  """Agent that streams nothing and answers with one final AI message."""
206
 
 
955
  )
956
  self.assertTrue(build_agent_mock.call_args.kwargs["include_thoughts"])
957
 
958
+ def test_stream_chat_publishes_completed_tool_args_before_result(self) -> None:
959
+ agent = FakeIncrementalRetrievalAgent([])
960
+ self.addCleanup(_drop_thread_record, "thread_incremental_tool")
961
+ request = ChatRequest(
962
+ query="Where should I store API keys?",
963
+ source_keys=("langchain",),
964
+ model_name=DEEPSEEK_DIRECT_MODEL_NAME,
965
+ include_reasoning=False,
966
+ enabled_tools=(),
967
+ )
968
+
969
+ async def collect_events():
970
+ return [event async for event in stream_chat(request)]
971
+
972
+ with (
973
+ patch("app.chat_service.build_agent", return_value=agent),
974
+ patch(
975
+ "app.chat_service.new_thread_id",
976
+ return_value="thread_incremental_tool",
977
+ ),
978
+ ):
979
+ events = asyncio.run(collect_events())
980
+
981
+ event_types = [event.type for event in events]
982
+ self.assertLess(
983
+ event_types.index("tool_call_started"),
984
+ event_types.index("tool_call_args_available"),
985
+ )
986
+ self.assertLess(
987
+ event_types.index("tool_call_args_available"),
988
+ event_types.index("tool_call_completed"),
989
+ )
990
+ args_event = next(
991
+ event for event in events if event.type == "tool_call_args_available"
992
+ )
993
+ self.assertEqual(
994
+ args_event.data["args"],
995
+ {"query": "secure API key storage"},
996
+ )
997
+ completed = next(
998
+ event for event in events if event.type == "tool_call_completed"
999
+ )
1000
+ self.assertEqual(len(completed.data["matches"]), 1)
1001
+ self.assertEqual(completed.data["matches"][0]["doc_id"], "doc-keys")
1002
+
1003
  def test_stream_chat_resolves_shell_citation_after_final_answer(self) -> None:
1004
  agent = FakeStreamingAgent([])
1005
  self.addCleanup(_drop_thread_record, "thread_rg")
tests/test_chroma_rag.py CHANGED
@@ -5,7 +5,7 @@ import pickle
5
  import tempfile
6
  import unittest
7
  from pathlib import Path
8
- from unittest.mock import patch
9
 
10
  import chromadb
11
  import tiktoken
@@ -351,11 +351,123 @@ class TokenBudgetTestCase(unittest.TestCase):
351
  self.assertEqual(len(kept_all), 3)
352
 
353
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
354
  class CollectionOpenTestCase(unittest.TestCase):
355
- def _write_document_dict(self, directory: str) -> str:
 
 
 
 
356
  path = Path(directory) / "document_dict_test.pkl"
357
  with open(path, "wb") as handle:
358
- pickle.dump({}, handle)
359
  return str(path)
360
 
361
  def test_init_fails_loudly_when_collection_missing(self) -> None:
@@ -383,7 +495,13 @@ class CollectionOpenTestCase(unittest.TestCase):
383
  chromadb.PersistentClient(path=temp_dir).create_collection(
384
  name="test-collection"
385
  )
386
- document_dict_path = self._write_document_dict(temp_dir)
 
 
 
 
 
 
387
 
388
  retriever = LocalChromaRetriever(
389
  db_path=temp_dir,
@@ -393,6 +511,10 @@ class CollectionOpenTestCase(unittest.TestCase):
393
  )
394
 
395
  self.assertEqual(retriever._collection.name, "test-collection")
 
 
 
 
396
 
397
 
398
  if __name__ == "__main__":
 
5
  import tempfile
6
  import unittest
7
  from pathlib import Path
8
+ from unittest.mock import ANY, Mock, patch
9
 
10
  import chromadb
11
  import tiktoken
 
351
  self.assertEqual(len(kept_all), 3)
352
 
353
 
354
+ class SourceNormalizationAndTimingTestCase(unittest.TestCase):
355
+ def _retriever(self) -> LocalChromaRetriever:
356
+ retriever = LocalChromaRetriever.__new__(LocalChromaRetriever)
357
+ retriever._indexed_sources = frozenset({"langchain", "peft"})
358
+ retriever._rrf_k = 60
359
+ retriever._bm25_top_k = 30
360
+ retriever._fusion_top_k = 30
361
+ retriever._rerank_model = "rerank-test"
362
+ retriever._rerank_top_k = 5
363
+ retriever._token_budget = 100_000
364
+ retriever._cohere = object()
365
+ return retriever
366
+
367
+ def _result(self) -> SearchResult:
368
+ return SearchResult(
369
+ chunk_id="chunk-1",
370
+ doc_id="doc-1",
371
+ title="Doc",
372
+ url="https://example.com",
373
+ source="peft",
374
+ retrieve_doc=False,
375
+ tokens=10,
376
+ score=0.9,
377
+ content="content",
378
+ chunk_content="content",
379
+ retrieval_method="dense",
380
+ )
381
+
382
+ def test_complete_indexed_source_selection_omits_filter(self) -> None:
383
+ retriever = self._retriever()
384
+
385
+ effective, mode = retriever._effective_allowed_sources(["langchain", "peft"])
386
+ self.assertIsNone(effective)
387
+ self.assertEqual(mode, "all_sources_omitted")
388
+
389
+ # Unknown requested keys do not change the result set, so a superset
390
+ # still safely covers every source present in the retrieval artifacts.
391
+ effective, mode = retriever._effective_allowed_sources(
392
+ ["langchain", "peft", "future-source"]
393
+ )
394
+ self.assertIsNone(effective)
395
+ self.assertEqual(mode, "all_sources_omitted")
396
+
397
+ def test_partial_source_selection_keeps_filter(self) -> None:
398
+ retriever = self._retriever()
399
+
400
+ sources = ["peft"]
401
+ effective, mode = retriever._effective_allowed_sources(sources)
402
+
403
+ self.assertEqual(effective, sources)
404
+ self.assertEqual(mode, "single_source")
405
+
406
+ def test_search_logs_and_traces_stage_timings_after_normalization(self) -> None:
407
+ retriever = self._retriever()
408
+ result = self._result()
409
+ retriever._dense_search = Mock(return_value=[result])
410
+ retriever._bm25_search = Mock(return_value=[])
411
+ retriever._apply_token_budget = Mock(return_value=[result])
412
+
413
+ class _FakeRun:
414
+ def __init__(self) -> None:
415
+ self.metadata: dict[str, object] = {}
416
+
417
+ def add_metadata(self, metadata: dict[str, object]) -> None:
418
+ self.metadata.update(metadata)
419
+
420
+ fake_run = _FakeRun()
421
+ with (
422
+ patch(
423
+ "app.chroma_rag.rerank_results",
424
+ return_value=[result],
425
+ ),
426
+ patch(
427
+ "app.chroma_rag.get_current_run_tree",
428
+ return_value=fake_run,
429
+ ),
430
+ self.assertLogs("app.chroma_rag", level="INFO") as captured_logs,
431
+ ):
432
+ results = retriever.search(
433
+ "adapter configuration",
434
+ allowed_sources=["langchain", "peft"],
435
+ )
436
+
437
+ self.assertEqual(results, [result])
438
+ retriever._dense_search.assert_called_once_with(
439
+ "adapter configuration",
440
+ allowed_sources=None,
441
+ timings=ANY,
442
+ )
443
+ retriever._bm25_search.assert_called_once_with(
444
+ "adapter configuration",
445
+ allowed_sources=None,
446
+ )
447
+ timing = fake_run.metadata["retrieval_timing"]
448
+ assert isinstance(timing, dict)
449
+ self.assertEqual(timing["filter_mode"], "all_sources_omitted")
450
+ self.assertEqual(timing["requested_source_count"], 2)
451
+ self.assertEqual(timing["indexed_source_count"], 2)
452
+ self.assertEqual(timing["result_count"], 1)
453
+ self.assertIn("stage_ms", timing)
454
+ self.assertTrue(
455
+ any(
456
+ "filter_mode=all_sources_omitted" in message and "chroma_ms=" in message
457
+ for message in captured_logs.output
458
+ )
459
+ )
460
+
461
+
462
  class CollectionOpenTestCase(unittest.TestCase):
463
+ def _write_document_dict(
464
+ self,
465
+ directory: str,
466
+ documents: dict[str, dict[str, object]] | None = None,
467
+ ) -> str:
468
  path = Path(directory) / "document_dict_test.pkl"
469
  with open(path, "wb") as handle:
470
+ pickle.dump(documents or {}, handle)
471
  return str(path)
472
 
473
  def test_init_fails_loudly_when_collection_missing(self) -> None:
 
495
  chromadb.PersistentClient(path=temp_dir).create_collection(
496
  name="test-collection"
497
  )
498
+ document_dict_path = self._write_document_dict(
499
+ temp_dir,
500
+ {
501
+ "doc-1": {"source": "peft"},
502
+ "doc-2": {"source": "langchain"},
503
+ },
504
+ )
505
 
506
  retriever = LocalChromaRetriever(
507
  db_path=temp_dir,
 
511
  )
512
 
513
  self.assertEqual(retriever._collection.name, "test-collection")
514
+ self.assertEqual(
515
+ retriever._indexed_sources,
516
+ frozenset({"peft", "langchain"}),
517
+ )
518
 
519
 
520
  if __name__ == "__main__":
tests/test_memory_variants.py CHANGED
@@ -9,6 +9,7 @@ no model client, no API keys, no vector DB.
9
  from __future__ import annotations
10
 
11
  import hashlib
 
12
  import unittest
13
  from types import SimpleNamespace
14
  from unittest import mock
@@ -300,6 +301,43 @@ class StableToolOutputCapTests(unittest.TestCase):
300
  signals["tool_output_retained_bytes"],
301
  )
302
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
303
 
304
  class ExperimentCompactionMiddlewareTests(unittest.TestCase):
305
  class FakeModel:
 
9
  from __future__ import annotations
10
 
11
  import hashlib
12
+ import json
13
  import unittest
14
  from types import SimpleNamespace
15
  from unittest import mock
 
301
  signals["tool_output_retained_bytes"],
302
  )
303
 
304
+ def test_cap_preserves_retrieval_chunk_metadata(self) -> None:
305
+ raw = json.dumps(
306
+ {
307
+ "query": "secure API key storage",
308
+ "matches": [
309
+ {
310
+ "chunk_id": "chunk-1",
311
+ "doc_id": "doc-1",
312
+ "title": "Manage API keys",
313
+ "url": "https://example.com/keys",
314
+ "source": "langchain",
315
+ "retrieve_doc": True,
316
+ "tokens": 1200,
317
+ "score": 0.9,
318
+ "content": "x" * 5_000,
319
+ "chunk_content": "API key guidance",
320
+ "heading_path": "Testing > Manage API keys",
321
+ "retrieval_method": "hybrid",
322
+ }
323
+ ],
324
+ }
325
+ )
326
+ result = StableToolOutputCapMiddleware(2_048)._cap(
327
+ make_request([], "retrieval-cap"),
328
+ ToolMessage(
329
+ content=raw,
330
+ name="retrieve_tutor_context",
331
+ tool_call_id="call-retrieval",
332
+ ),
333
+ )
334
+
335
+ self.assertIn("tool output truncated", result.content)
336
+ retained = result.additional_kwargs["stable_tool_cap"]["retrieval_matches"]
337
+ self.assertEqual(len(retained), 1)
338
+ self.assertEqual(retained[0]["doc_id"], "doc-1")
339
+ self.assertEqual(retained[0]["source_key"], "langchain")
340
+
341
 
342
  class ExperimentCompactionMiddlewareTests(unittest.TestCase):
343
  class FakeModel: