omarsol commited on
Commit
3f613d2
·
1 Parent(s): 4f62444

ruff format

Browse files
data/scraping_scripts/add_context_to_nodes.py CHANGED
@@ -19,19 +19,19 @@ from pydantic import BaseModel, Field
19
  from tenacity import retry, retry_if_exception, stop_after_attempt
20
  from tqdm.asyncio import tqdm
21
 
22
- from scripts.chroma_rag import ChunkRecord, build_chunk_records, format_chunk_for_retrieval
 
 
 
 
23
 
24
  load_dotenv(".env")
25
 
26
  CONTEXT_MODEL = os.getenv("GEMINI_CONTEXT_MODEL", "gemini-3.1-flash-lite")
27
  CONTEXT_MAX_OUTPUT_TOKENS = 1000
28
  CONTEXT_TPM_LIMIT = int(os.getenv("GEMINI_CONTEXT_TPM_LIMIT", "30000000"))
29
- CONTEXT_TPM_SAFETY_MARGIN = float(
30
- os.getenv("GEMINI_CONTEXT_TPM_SAFETY_MARGIN", "0.8")
31
- )
32
- CONTEXT_TPM_WINDOW_SECONDS = float(
33
- os.getenv("GEMINI_CONTEXT_TPM_WINDOW_SECONDS", "60")
34
- )
35
  CONTEXT_RETRY_ATTEMPTS = int(os.getenv("GEMINI_CONTEXT_RETRY_ATTEMPTS", "8"))
36
  DEFAULT_SEMAPHORE_LIMIT = int(os.getenv("GEMINI_CONTEXT_CONCURRENCY", "50"))
37
  MAX_DOCUMENT_TOKENS = 120_000
@@ -192,7 +192,7 @@ def wait_for_genai_retry(retry_state) -> float:
192
  if retry_delay is not None:
193
  return retry_delay + random.uniform(1.0, 4.0)
194
 
195
- exponential_delay = min(60, max(4, 2 ** retry_state.attempt_number))
196
  return exponential_delay + random.uniform(0.5, 2.0)
197
 
198
 
 
19
  from tenacity import retry, retry_if_exception, stop_after_attempt
20
  from tqdm.asyncio import tqdm
21
 
22
+ from scripts.chroma_rag import (
23
+ ChunkRecord,
24
+ build_chunk_records,
25
+ format_chunk_for_retrieval,
26
+ )
27
 
28
  load_dotenv(".env")
29
 
30
  CONTEXT_MODEL = os.getenv("GEMINI_CONTEXT_MODEL", "gemini-3.1-flash-lite")
31
  CONTEXT_MAX_OUTPUT_TOKENS = 1000
32
  CONTEXT_TPM_LIMIT = int(os.getenv("GEMINI_CONTEXT_TPM_LIMIT", "30000000"))
33
+ CONTEXT_TPM_SAFETY_MARGIN = float(os.getenv("GEMINI_CONTEXT_TPM_SAFETY_MARGIN", "0.8"))
34
+ CONTEXT_TPM_WINDOW_SECONDS = float(os.getenv("GEMINI_CONTEXT_TPM_WINDOW_SECONDS", "60"))
 
 
 
 
35
  CONTEXT_RETRY_ATTEMPTS = int(os.getenv("GEMINI_CONTEXT_RETRY_ATTEMPTS", "8"))
36
  DEFAULT_SEMAPHORE_LIMIT = int(os.getenv("GEMINI_CONTEXT_CONCURRENCY", "50"))
37
  MAX_DOCUMENT_TOKENS = 120_000
 
192
  if retry_delay is not None:
193
  return retry_delay + random.uniform(1.0, 4.0)
194
 
195
+ exponential_delay = min(60, max(4, 2**retry_state.attempt_number))
196
  return exponential_delay + random.uniform(0.5, 2.0)
197
 
198
 
data/scraping_scripts/build_kb_artifacts.py CHANGED
@@ -217,9 +217,7 @@ def unique_path(candidate: Path, used_paths: set[Path], content_hash: str) -> Pa
217
  candidate = candidate.with_name(f"{stem}-{suffix}{candidate.suffix}")
218
  counter = 2
219
  while candidate in used_paths:
220
- candidate = candidate.with_name(
221
- f"{stem}-{suffix}-{counter}{candidate.suffix}"
222
- )
223
  counter += 1
224
  used_paths.add(candidate)
225
  return candidate
@@ -373,8 +371,11 @@ def write_raw_markdown(
373
  for record in records:
374
  original = find_original_markdown(record, original_indexes)
375
  if original:
376
- output_path = raw_dir / "docs" / record.source / safe_relative_path(
377
- original.source_path
 
 
 
378
  )
379
  output_path = unique_path(output_path, used_paths, record.content_hash)
380
  source_path = original.source_path
@@ -386,7 +387,9 @@ def write_raw_markdown(
386
  original_path = ""
387
  content = record.content
388
  if record.source_group == "docs":
389
- fallback_counts[record.source] = fallback_counts.get(record.source, 0) + 1
 
 
390
 
391
  output_path.parent.mkdir(parents=True, exist_ok=True)
392
  markdown = markdown_with_frontmatter(
@@ -417,7 +420,9 @@ def write_raw_markdown(
417
  def iter_markdown_headings(path: Path) -> Iterable[dict[str, Any]]:
418
  in_fence = False
419
  fence_char = ""
420
- for line_number, line in enumerate(path.read_text(encoding="utf-8").splitlines(), 1):
 
 
421
  fence_match = FENCE_RE.match(line)
422
  if fence_match:
423
  current = fence_match.group(1)[0]
@@ -464,7 +469,9 @@ def extract_symbols(text: str) -> set[str]:
464
  }
465
 
466
 
467
- def write_generated_indexes(manifest: list[dict[str, Any]], generated_dir: Path) -> None:
 
 
468
  if generated_dir.exists():
469
  shutil.rmtree(generated_dir)
470
  generated_dir.mkdir(parents=True, exist_ok=True)
@@ -511,7 +518,9 @@ def write_generated_indexes(manifest: list[dict[str, Any]], generated_dir: Path)
511
  for row in heading_rows:
512
  handle.write(json.dumps(row, ensure_ascii=False, sort_keys=True) + "\n")
513
 
514
- with (generated_dir / "symbols.tsv").open("w", encoding="utf-8", newline="") as handle:
 
 
515
  writer = csv.DictWriter(
516
  handle,
517
  fieldnames=["symbol", "source", "title", "path", "heading", "doc_id"],
@@ -523,7 +532,9 @@ def write_generated_indexes(manifest: list[dict[str, Any]], generated_dir: Path)
523
 
524
  def build_kb_artifacts(input_file: Path, output_dir: Path) -> dict[str, int]:
525
  rows = load_jsonl(input_file)
526
- records = [normalize_record(row) for row in rows if str(row.get("content") or "").strip()]
 
 
527
  manifest = write_raw_markdown(records, input_file=input_file, output_dir=output_dir)
528
  write_generated_indexes(manifest, output_dir / GENERATED_DIR_NAME)
529
  return {"documents": len(records), "manifest_rows": len(manifest)}
 
217
  candidate = candidate.with_name(f"{stem}-{suffix}{candidate.suffix}")
218
  counter = 2
219
  while candidate in used_paths:
220
+ candidate = candidate.with_name(f"{stem}-{suffix}-{counter}{candidate.suffix}")
 
 
221
  counter += 1
222
  used_paths.add(candidate)
223
  return candidate
 
371
  for record in records:
372
  original = find_original_markdown(record, original_indexes)
373
  if original:
374
+ output_path = (
375
+ raw_dir
376
+ / "docs"
377
+ / record.source
378
+ / safe_relative_path(original.source_path)
379
  )
380
  output_path = unique_path(output_path, used_paths, record.content_hash)
381
  source_path = original.source_path
 
387
  original_path = ""
388
  content = record.content
389
  if record.source_group == "docs":
390
+ fallback_counts[record.source] = (
391
+ fallback_counts.get(record.source, 0) + 1
392
+ )
393
 
394
  output_path.parent.mkdir(parents=True, exist_ok=True)
395
  markdown = markdown_with_frontmatter(
 
420
  def iter_markdown_headings(path: Path) -> Iterable[dict[str, Any]]:
421
  in_fence = False
422
  fence_char = ""
423
+ for line_number, line in enumerate(
424
+ path.read_text(encoding="utf-8").splitlines(), 1
425
+ ):
426
  fence_match = FENCE_RE.match(line)
427
  if fence_match:
428
  current = fence_match.group(1)[0]
 
469
  }
470
 
471
 
472
+ def write_generated_indexes(
473
+ manifest: list[dict[str, Any]], generated_dir: Path
474
+ ) -> None:
475
  if generated_dir.exists():
476
  shutil.rmtree(generated_dir)
477
  generated_dir.mkdir(parents=True, exist_ok=True)
 
518
  for row in heading_rows:
519
  handle.write(json.dumps(row, ensure_ascii=False, sort_keys=True) + "\n")
520
 
521
+ with (generated_dir / "symbols.tsv").open(
522
+ "w", encoding="utf-8", newline=""
523
+ ) as handle:
524
  writer = csv.DictWriter(
525
  handle,
526
  fieldnames=["symbol", "source", "title", "path", "heading", "doc_id"],
 
532
 
533
  def build_kb_artifacts(input_file: Path, output_dir: Path) -> dict[str, int]:
534
  rows = load_jsonl(input_file)
535
+ records = [
536
+ normalize_record(row) for row in rows if str(row.get("content") or "").strip()
537
+ ]
538
  manifest = write_raw_markdown(records, input_file=input_file, output_dir=output_dir)
539
  write_generated_indexes(manifest, output_dir / GENERATED_DIR_NAME)
540
  return {"documents": len(records), "manifest_rows": len(manifest)}
data/scraping_scripts/create_vector_stores.py CHANGED
@@ -77,7 +77,9 @@ SOURCE_CONFIGS = vector_store_source_configs()
77
 
78
  def load_or_create_chunk_records(source: str) -> list[Any]:
79
  config = SOURCE_CONFIGS[source]
80
- if source == "all_sources" and os.path.exists("data/all_sources_contextual_nodes.pkl"):
 
 
81
  with open("data/all_sources_contextual_nodes.pkl", "rb") as handle:
82
  records = pickle.load(handle)
83
  active_records = []
 
77
 
78
  def load_or_create_chunk_records(source: str) -> list[Any]:
79
  config = SOURCE_CONFIGS[source]
80
+ if source == "all_sources" and os.path.exists(
81
+ "data/all_sources_contextual_nodes.pkl"
82
+ ):
83
  with open("data/all_sources_contextual_nodes.pkl", "rb") as handle:
84
  records = pickle.load(handle)
85
  active_records = []
data/scraping_scripts/github_to_markdown_ai_docs.py CHANGED
@@ -22,7 +22,7 @@ Example:
22
 
23
  This will download and process the documentation files for both TRL and PEFT libraries.
24
 
25
- Note:
26
  - Ensure you have set the GITHUB_TOKEN variable with your GitHub Personal Access Token.
27
  - The script creates a 'data' directory in the current working directory to store the downloaded files.
28
  - Each source's files are stored in a subdirectory named '<repo>_md_files'.
@@ -363,9 +363,7 @@ def process_source(source: str):
363
  os.makedirs(target_dir, exist_ok=True)
364
  fetch_files(api_url, target_dir, source_extensions, local_dir)
365
  else:
366
- api_url = (
367
- f"https://api.github.com/repos/{config['owner']}/{config['repo']}/contents/{config['path']}"
368
- )
369
  fetch_files(api_url, local_dir, source_extensions, local_dir)
370
 
371
  save_source_extension_manifest(local_dir, source_extensions)
 
22
 
23
  This will download and process the documentation files for both TRL and PEFT libraries.
24
 
25
+ Note:
26
  - Ensure you have set the GITHUB_TOKEN variable with your GitHub Personal Access Token.
27
  - The script creates a 'data' directory in the current working directory to store the downloaded files.
28
  - Each source's files are stored in a subdirectory named '<repo>_md_files'.
 
363
  os.makedirs(target_dir, exist_ok=True)
364
  fetch_files(api_url, target_dir, source_extensions, local_dir)
365
  else:
366
+ api_url = f"https://api.github.com/repos/{config['owner']}/{config['repo']}/contents/{config['path']}"
 
 
367
  fetch_files(api_url, local_dir, source_extensions, local_dir)
368
 
369
  save_source_extension_manifest(local_dir, source_extensions)
data/scraping_scripts/retire_source_workflow.py CHANGED
@@ -85,7 +85,9 @@ def ensure_required_files_exist(data_repo_id: str) -> None:
85
  if local_path.exists():
86
  continue
87
 
88
- print(f"{remote_filename} not found locally. Downloading from {data_repo_id}...")
 
 
89
  hf_hub_download(
90
  token=os.getenv("HF_TOKEN"),
91
  repo_id=data_repo_id,
@@ -95,7 +97,9 @@ def ensure_required_files_exist(data_repo_id: str) -> None:
95
  )
96
 
97
 
98
- def filter_all_sources_jsonl(sources_to_retire: set[str], dry_run: bool) -> Counter[str]:
 
 
99
  if not ALL_SOURCES_JSONL.exists():
100
  raise SystemExit(f"Missing required file: {ALL_SOURCES_JSONL}")
101
 
@@ -272,7 +276,9 @@ def update_source_registry(sources: list[str], dry_run: bool) -> None:
272
 
273
  def rebuild_vector_store() -> None:
274
  if not os.getenv("COHERE_API_KEY"):
275
- raise SystemExit("COHERE_API_KEY is required to rebuild the Chroma vector store.")
 
 
276
 
277
  print("Rebuilding Chroma vector store for all_sources...")
278
  result = run_module("data.scraping_scripts.create_vector_stores", "all_sources")
@@ -318,7 +324,9 @@ def delete_remote_source_files(
318
 
319
  def upload_vector_store(vector_repo_id: str) -> None:
320
  print(f"Uploading rebuilt vector store to {vector_repo_id}...")
321
- result = run_module("data.scraping_scripts.upload_dbs_to_hf", "--repo", vector_repo_id)
 
 
322
  if result.returncode != 0:
323
  raise SystemExit("Error uploading vector store. Check output above.")
324
 
@@ -406,9 +414,8 @@ def main() -> None:
406
  missing_required_files = [
407
  path for path in (ALL_SOURCES_JSONL, CONTEXTUAL_NODES) if not path.exists()
408
  ]
409
- needs_data_repo = (
410
- (not args.skip_download and bool(missing_required_files))
411
- or (not args.dry_run and not args.skip_upload and not args.skip_data_upload)
412
  )
413
  needs_vector_repo = (
414
  not args.dry_run and not args.skip_upload and not args.skip_vector_rebuild
 
85
  if local_path.exists():
86
  continue
87
 
88
+ print(
89
+ f"{remote_filename} not found locally. Downloading from {data_repo_id}..."
90
+ )
91
  hf_hub_download(
92
  token=os.getenv("HF_TOKEN"),
93
  repo_id=data_repo_id,
 
97
  )
98
 
99
 
100
+ def filter_all_sources_jsonl(
101
+ sources_to_retire: set[str], dry_run: bool
102
+ ) -> Counter[str]:
103
  if not ALL_SOURCES_JSONL.exists():
104
  raise SystemExit(f"Missing required file: {ALL_SOURCES_JSONL}")
105
 
 
276
 
277
  def rebuild_vector_store() -> None:
278
  if not os.getenv("COHERE_API_KEY"):
279
+ raise SystemExit(
280
+ "COHERE_API_KEY is required to rebuild the Chroma vector store."
281
+ )
282
 
283
  print("Rebuilding Chroma vector store for all_sources...")
284
  result = run_module("data.scraping_scripts.create_vector_stores", "all_sources")
 
324
 
325
  def upload_vector_store(vector_repo_id: str) -> None:
326
  print(f"Uploading rebuilt vector store to {vector_repo_id}...")
327
+ result = run_module(
328
+ "data.scraping_scripts.upload_dbs_to_hf", "--repo", vector_repo_id
329
+ )
330
  if result.returncode != 0:
331
  raise SystemExit("Error uploading vector store. Check output above.")
332
 
 
414
  missing_required_files = [
415
  path for path in (ALL_SOURCES_JSONL, CONTEXTUAL_NODES) if not path.exists()
416
  ]
417
+ needs_data_repo = (not args.skip_download and bool(missing_required_files)) or (
418
+ not args.dry_run and not args.skip_upload and not args.skip_data_upload
 
419
  )
420
  needs_vector_repo = (
421
  not args.dry_run and not args.skip_upload and not args.skip_vector_rebuild
data/scraping_scripts/update_kb_wiki.py CHANGED
@@ -64,7 +64,9 @@ def generated_block(content: str) -> str:
64
  return f"{AUTO_START}\n{content.rstrip()}\n{AUTO_END}"
65
 
66
 
67
- def write_generated_section(path: Path, header: str, generated: str, *, overwrite: bool) -> None:
 
 
68
  path.parent.mkdir(parents=True, exist_ok=True)
69
  block = generated_block(generated)
70
  if not path.exists() or overwrite:
@@ -86,7 +88,9 @@ def top_sources(manifest: list[dict[str, Any]]) -> list[tuple[str, int]]:
86
  return sorted(counts.items(), key=lambda item: (-item[1], item[0]))
87
 
88
 
89
- def seed_index(kb_dir: Path, manifest: list[dict[str, Any]], *, overwrite: bool) -> None:
 
 
90
  group_by_source = {
91
  str(row.get("source") or "unknown"): str(row.get("source_group") or "docs")
92
  for row in manifest
@@ -94,7 +98,9 @@ def seed_index(kb_dir: Path, manifest: list[dict[str, Any]], *, overwrite: bool)
94
  source_lines = []
95
  for source, count in top_sources(manifest):
96
  folder = "courses" if group_by_source.get(source) == "courses" else "frameworks"
97
- source_lines.append(f"- {source}: `wiki/{folder}/{source}.md` - {count} corpus pages")
 
 
98
  topic_lines = [f"- {topic}: `wiki/topics/{topic}.md`" for topic in TOPIC_KEYWORDS]
99
  header = """# AI Tutor KB Index
100
 
@@ -136,7 +142,9 @@ def seed_log(kb_dir: Path, manifest: list[dict[str, Any]], *, overwrite: bool) -
136
  path.write_text(f"{existing}\n\n{entry.rstrip()}\n", encoding="utf-8")
137
 
138
 
139
- def seed_source_pages(kb_dir: Path, manifest: list[dict[str, Any]], *, overwrite: bool) -> None:
 
 
140
  by_source: dict[str, list[dict[str, Any]]] = defaultdict(list)
141
  for row in manifest:
142
  by_source[str(row.get("source") or "unknown")].append(row)
@@ -145,8 +153,7 @@ def seed_source_pages(kb_dir: Path, manifest: list[dict[str, Any]], *, overwrite
145
  rows = sorted(rows, key=lambda item: str(item.get("title") or ""))
146
  sample = rows[:20]
147
  page_links = [
148
- f"- {row.get('title')}: `{shell_path(row, kb_dir)}`"
149
- for row in sample
150
  ]
151
  group = str(rows[0].get("source_group") or "docs")
152
  folder = "courses" if group == "courses" else "frameworks"
@@ -190,13 +197,12 @@ def matching_topic_rows(
190
  return [row for _score, row in scored[:limit]]
191
 
192
 
193
- def seed_topic_pages(kb_dir: Path, manifest: list[dict[str, Any]], *, overwrite: bool) -> None:
 
 
194
  for topic, keywords in TOPIC_KEYWORDS.items():
195
  rows = matching_topic_rows(manifest, keywords)
196
- links = [
197
- f"- {row.get('title')}: `{shell_path(row, kb_dir)}`"
198
- for row in rows
199
- ]
200
  header = f"""# {topic.replace("-", " ").title()}
201
 
202
  Use this page as a starting map for questions involving: {", ".join(keywords)}.
 
64
  return f"{AUTO_START}\n{content.rstrip()}\n{AUTO_END}"
65
 
66
 
67
+ def write_generated_section(
68
+ path: Path, header: str, generated: str, *, overwrite: bool
69
+ ) -> None:
70
  path.parent.mkdir(parents=True, exist_ok=True)
71
  block = generated_block(generated)
72
  if not path.exists() or overwrite:
 
88
  return sorted(counts.items(), key=lambda item: (-item[1], item[0]))
89
 
90
 
91
+ def seed_index(
92
+ kb_dir: Path, manifest: list[dict[str, Any]], *, overwrite: bool
93
+ ) -> None:
94
  group_by_source = {
95
  str(row.get("source") or "unknown"): str(row.get("source_group") or "docs")
96
  for row in manifest
 
98
  source_lines = []
99
  for source, count in top_sources(manifest):
100
  folder = "courses" if group_by_source.get(source) == "courses" else "frameworks"
101
+ source_lines.append(
102
+ f"- {source}: `wiki/{folder}/{source}.md` - {count} corpus pages"
103
+ )
104
  topic_lines = [f"- {topic}: `wiki/topics/{topic}.md`" for topic in TOPIC_KEYWORDS]
105
  header = """# AI Tutor KB Index
106
 
 
142
  path.write_text(f"{existing}\n\n{entry.rstrip()}\n", encoding="utf-8")
143
 
144
 
145
+ def seed_source_pages(
146
+ kb_dir: Path, manifest: list[dict[str, Any]], *, overwrite: bool
147
+ ) -> None:
148
  by_source: dict[str, list[dict[str, Any]]] = defaultdict(list)
149
  for row in manifest:
150
  by_source[str(row.get("source") or "unknown")].append(row)
 
153
  rows = sorted(rows, key=lambda item: str(item.get("title") or ""))
154
  sample = rows[:20]
155
  page_links = [
156
+ f"- {row.get('title')}: `{shell_path(row, kb_dir)}`" for row in sample
 
157
  ]
158
  group = str(rows[0].get("source_group") or "docs")
159
  folder = "courses" if group == "courses" else "frameworks"
 
197
  return [row for _score, row in scored[:limit]]
198
 
199
 
200
+ def seed_topic_pages(
201
+ kb_dir: Path, manifest: list[dict[str, Any]], *, overwrite: bool
202
+ ) -> None:
203
  for topic, keywords in TOPIC_KEYWORDS.items():
204
  rows = matching_topic_rows(manifest, keywords)
205
+ links = [f"- {row.get('title')}: `{shell_path(row, kb_dir)}`" for row in rows]
 
 
 
206
  header = f"""# {topic.replace("-", " ").title()}
207
 
208
  Use this page as a starting map for questions involving: {", ".join(keywords)}.
scripts/api.py CHANGED
@@ -95,7 +95,9 @@ app.add_middleware(
95
 
96
 
97
  def sse_frame(payload: dict[str, Any] | str) -> str:
98
- data = payload if isinstance(payload, str) else json.dumps(payload, ensure_ascii=False)
 
 
99
  return f"data: {data}\n\n"
100
 
101
 
@@ -138,11 +140,14 @@ def build_chat_request(payload: ApiChatRequest) -> ChatRequest:
138
 
139
  allowed_source_keys = set(AVAILABLE_SOURCES)
140
  requested_source_keys = payload.sourceKeys or list(DEFAULT_SELECTED_SOURCE_KEYS)
141
- source_keys = tuple(
142
- dict.fromkeys(
143
- key for key in requested_source_keys if key in allowed_source_keys
 
 
144
  )
145
- ) or DEFAULT_SELECTED_SOURCE_KEYS
 
146
  model_name = (payload.model or DEFAULT_MODEL_NAME).strip()
147
  if payload.enabledTools is None:
148
  enabled_tools = tuple(
@@ -226,7 +231,9 @@ class UIMessageStreamEncoder:
226
  if event.type == "reasoning_delta":
227
  if not self.active_reasoning_id:
228
  self.active_reasoning_id = f"reasoning_{uuid4().hex[:8]}"
229
- parts.append({"type": "reasoning-start", "id": self.active_reasoning_id})
 
 
230
  parts.append(
231
  {
232
  "type": "reasoning-delta",
@@ -281,7 +288,9 @@ class UIMessageStreamEncoder:
281
  "group": str(event.data.get("group", "")),
282
  }
283
  if call_id:
284
- self.source_matches_by_call_id.setdefault(call_id, []).append(source_data)
 
 
285
  parts.append(
286
  {
287
  "type": "source-url",
@@ -328,7 +337,7 @@ class UIMessageStreamEncoder:
328
  "id": self.text_block_id,
329
  "delta": answer,
330
  }
331
- )
332
  if self.text_block_id:
333
  parts.append({"type": "text-end", "id": self.text_block_id})
334
  parts.append({"type": "finish-step"})
 
95
 
96
 
97
  def sse_frame(payload: dict[str, Any] | str) -> str:
98
+ data = (
99
+ payload if isinstance(payload, str) else json.dumps(payload, ensure_ascii=False)
100
+ )
101
  return f"data: {data}\n\n"
102
 
103
 
 
140
 
141
  allowed_source_keys = set(AVAILABLE_SOURCES)
142
  requested_source_keys = payload.sourceKeys or list(DEFAULT_SELECTED_SOURCE_KEYS)
143
+ source_keys = (
144
+ tuple(
145
+ dict.fromkeys(
146
+ key for key in requested_source_keys if key in allowed_source_keys
147
+ )
148
  )
149
+ or DEFAULT_SELECTED_SOURCE_KEYS
150
+ )
151
  model_name = (payload.model or DEFAULT_MODEL_NAME).strip()
152
  if payload.enabledTools is None:
153
  enabled_tools = tuple(
 
231
  if event.type == "reasoning_delta":
232
  if not self.active_reasoning_id:
233
  self.active_reasoning_id = f"reasoning_{uuid4().hex[:8]}"
234
+ parts.append(
235
+ {"type": "reasoning-start", "id": self.active_reasoning_id}
236
+ )
237
  parts.append(
238
  {
239
  "type": "reasoning-delta",
 
288
  "group": str(event.data.get("group", "")),
289
  }
290
  if call_id:
291
+ self.source_matches_by_call_id.setdefault(call_id, []).append(
292
+ source_data
293
+ )
294
  parts.append(
295
  {
296
  "type": "source-url",
 
337
  "id": self.text_block_id,
338
  "delta": answer,
339
  }
340
+ )
341
  if self.text_block_id:
342
  parts.append({"type": "text-end", "id": self.text_block_id})
343
  parts.append({"type": "finish-step"})
scripts/chat_service.py CHANGED
@@ -64,7 +64,9 @@ THOUGHTS_BLOCK_START = "<!-- GEMINI_THOUGHTS_START -->"
64
  THOUGHTS_BLOCK_END = "<!-- GEMINI_THOUGHTS_END -->"
65
  ANSWER_HEADER = "**Answer**"
66
  LEGACY_THOUGHTS_DETAILS_OPEN = "<details><summary>Gemini thoughts</summary>"
67
- LEGACY_THOUGHTS_DETAILS_OPEN_EXPANDED = "<details open><summary>Gemini thoughts</summary>"
 
 
68
  CHECKPOINTER = InMemorySaver()
69
  _RETRIEVER_INIT_LOCK = Lock()
70
  KB_TOOL_NAMES = ("run_kb_command",)
@@ -403,7 +405,9 @@ def collect_retrieval_source_matches(payload: str) -> list[SourceMatch]:
403
  return matches
404
 
405
 
406
- def _record_evidence(target: dict[str, SourceMatch], matches: list[SourceMatch]) -> None:
 
 
407
  for match in matches:
408
  key = source_match_key(match)
409
  existing = target.get(key)
@@ -704,9 +708,7 @@ class SourcePreferenceMiddleware(AgentMiddleware):
704
  label = SOURCE_KEY_TO_LABEL.get(key, key)
705
  group = "courses" if key in COURSE_SOURCE_KEYS else "docs"
706
  wiki_dir = "courses" if key in COURSE_SOURCE_KEYS else "frameworks"
707
- lines.append(
708
- f"- {label}: `raw/{group}/{key}/`, `wiki/{wiki_dir}/{key}.md`"
709
- )
710
  lines.append("")
711
  lines.append(
712
  "Only branch out to other KB sources if these don't have the answer."
@@ -1098,9 +1100,7 @@ async def stream_chat(request: ChatRequest) -> AsyncIterator[ChatEvent]:
1098
 
1099
  step = str(metadata.get("langgraph_step", ""))
1100
  if include_reasoning:
1101
- thought_text = "\n\n".join(
1102
- extract_thought_summaries(token.content)
1103
- )
1104
  if thought_text:
1105
  yield ChatEvent(
1106
  "reasoning_delta",
@@ -1157,9 +1157,7 @@ async def stream_chat(request: ChatRequest) -> AsyncIterator[ChatEvent]:
1157
  "call_id": google_search_call_id,
1158
  "tool_name": GOOGLE_SEARCH_TOOL_NAME,
1159
  "args": {
1160
- "query": "; ".join(new_queries)
1161
- if new_queries
1162
- else ""
1163
  },
1164
  "args_text": "; ".join(new_queries),
1165
  },
@@ -1245,9 +1243,7 @@ async def stream_chat(request: ChatRequest) -> AsyncIterator[ChatEvent]:
1245
  "call_id": google_search_call_id,
1246
  "tool_name": GOOGLE_SEARCH_TOOL_NAME,
1247
  "args": {
1248
- "query": "; ".join(new_queries)
1249
- if new_queries
1250
- else ""
1251
  },
1252
  "args_text": "; ".join(new_queries),
1253
  },
@@ -1311,9 +1307,7 @@ async def stream_chat(request: ChatRequest) -> AsyncIterator[ChatEvent]:
1311
  output_text = "Google search ran but returned no grounding results."
1312
  else:
1313
  plural = "" if google_search_match_count == 1 else "s"
1314
- output_text = (
1315
- f"Google search returned {google_search_match_count} web result{plural}."
1316
- )
1317
  yield ChatEvent(
1318
  "tool_call_completed",
1319
  {
 
64
  THOUGHTS_BLOCK_END = "<!-- GEMINI_THOUGHTS_END -->"
65
  ANSWER_HEADER = "**Answer**"
66
  LEGACY_THOUGHTS_DETAILS_OPEN = "<details><summary>Gemini thoughts</summary>"
67
+ LEGACY_THOUGHTS_DETAILS_OPEN_EXPANDED = (
68
+ "<details open><summary>Gemini thoughts</summary>"
69
+ )
70
  CHECKPOINTER = InMemorySaver()
71
  _RETRIEVER_INIT_LOCK = Lock()
72
  KB_TOOL_NAMES = ("run_kb_command",)
 
405
  return matches
406
 
407
 
408
+ def _record_evidence(
409
+ target: dict[str, SourceMatch], matches: list[SourceMatch]
410
+ ) -> None:
411
  for match in matches:
412
  key = source_match_key(match)
413
  existing = target.get(key)
 
708
  label = SOURCE_KEY_TO_LABEL.get(key, key)
709
  group = "courses" if key in COURSE_SOURCE_KEYS else "docs"
710
  wiki_dir = "courses" if key in COURSE_SOURCE_KEYS else "frameworks"
711
+ lines.append(f"- {label}: `raw/{group}/{key}/`, `wiki/{wiki_dir}/{key}.md`")
 
 
712
  lines.append("")
713
  lines.append(
714
  "Only branch out to other KB sources if these don't have the answer."
 
1100
 
1101
  step = str(metadata.get("langgraph_step", ""))
1102
  if include_reasoning:
1103
+ thought_text = "\n\n".join(extract_thought_summaries(token.content))
 
 
1104
  if thought_text:
1105
  yield ChatEvent(
1106
  "reasoning_delta",
 
1157
  "call_id": google_search_call_id,
1158
  "tool_name": GOOGLE_SEARCH_TOOL_NAME,
1159
  "args": {
1160
+ "query": "; ".join(new_queries) if new_queries else ""
 
 
1161
  },
1162
  "args_text": "; ".join(new_queries),
1163
  },
 
1243
  "call_id": google_search_call_id,
1244
  "tool_name": GOOGLE_SEARCH_TOOL_NAME,
1245
  "args": {
1246
+ "query": "; ".join(new_queries) if new_queries else ""
 
 
1247
  },
1248
  "args_text": "; ".join(new_queries),
1249
  },
 
1307
  output_text = "Google search ran but returned no grounding results."
1308
  else:
1309
  plural = "" if google_search_match_count == 1 else "s"
1310
+ output_text = f"Google search returned {google_search_match_count} web result{plural}."
 
 
1311
  yield ChatEvent(
1312
  "tool_call_completed",
1313
  {
scripts/chroma_rag.py CHANGED
@@ -515,8 +515,7 @@ def heading_aware_markdown_chunks(
515
  for split_unit in split_units:
516
  size = unit_size(split_unit.text)
517
  heading_changed = (
518
- current_heading_path
519
- and split_unit.heading_path != current_heading_path
520
  )
521
  would_exceed = current_parts and current_size + size > budget_size
522
  if heading_changed or would_exceed:
@@ -591,7 +590,7 @@ def build_chunk_records(
591
  }
592
  chunk_records.append(
593
  ChunkRecord(
594
- chunk_id=f'{document["doc_id"]}:{index}',
595
  doc_id=document["doc_id"],
596
  text=chunk.text,
597
  metadata=metadata,
@@ -836,7 +835,7 @@ def _wait_for_cohere_retry(
836
  if retry_after is not None:
837
  delay = retry_after
838
  else:
839
- delay = min(window_seconds, max(15.0, 2.0 ** attempt))
840
 
841
  time.sleep(delay + random.uniform(0.5, 2.0))
842
 
@@ -948,7 +947,10 @@ def embed_texts(
948
  )
949
  break
950
  except Exception as exc:
951
- if not _is_cohere_rate_limit_error(exc) or attempt == retry_attempts:
 
 
 
952
  raise
953
  _wait_for_cohere_retry(exc, attempt, window_seconds)
954
 
@@ -1096,7 +1098,9 @@ def reciprocal_rank_fusion(
1096
  current = representatives.get(key)
1097
  if current is None:
1098
  representatives[key] = result
1099
- elif result.retrieval_method == "bm25" and current.retrieval_method != "bm25":
 
 
1100
  representatives[key] = result
1101
  elif (
1102
  result.retrieval_method == current.retrieval_method
@@ -1357,7 +1361,9 @@ def ensure_parent_dir(path: str) -> None:
1357
  Path(path).parent.mkdir(parents=True, exist_ok=True)
1358
 
1359
 
1360
- def save_document_dict(document_dict: dict[str, dict[str, Any]], output_file: str) -> None:
 
 
1361
  ensure_parent_dir(output_file)
1362
  with open(output_file, "wb") as handle:
1363
  pickle.dump(document_dict, handle)
 
515
  for split_unit in split_units:
516
  size = unit_size(split_unit.text)
517
  heading_changed = (
518
+ current_heading_path and split_unit.heading_path != current_heading_path
 
519
  )
520
  would_exceed = current_parts and current_size + size > budget_size
521
  if heading_changed or would_exceed:
 
590
  }
591
  chunk_records.append(
592
  ChunkRecord(
593
+ chunk_id=f"{document['doc_id']}:{index}",
594
  doc_id=document["doc_id"],
595
  text=chunk.text,
596
  metadata=metadata,
 
835
  if retry_after is not None:
836
  delay = retry_after
837
  else:
838
+ delay = min(window_seconds, max(15.0, 2.0**attempt))
839
 
840
  time.sleep(delay + random.uniform(0.5, 2.0))
841
 
 
947
  )
948
  break
949
  except Exception as exc:
950
+ if (
951
+ not _is_cohere_rate_limit_error(exc)
952
+ or attempt == retry_attempts
953
+ ):
954
  raise
955
  _wait_for_cohere_retry(exc, attempt, window_seconds)
956
 
 
1098
  current = representatives.get(key)
1099
  if current is None:
1100
  representatives[key] = result
1101
+ elif (
1102
+ result.retrieval_method == "bm25" and current.retrieval_method != "bm25"
1103
+ ):
1104
  representatives[key] = result
1105
  elif (
1106
  result.retrieval_method == current.retrieval_method
 
1361
  Path(path).parent.mkdir(parents=True, exist_ok=True)
1362
 
1363
 
1364
+ def save_document_dict(
1365
+ document_dict: dict[str, dict[str, Any]], output_file: str
1366
+ ) -> None:
1367
  ensure_parent_dir(output_file)
1368
  with open(output_file, "wb") as handle:
1369
  pickle.dump(document_dict, handle)
scripts/gradio_presenter.py CHANGED
@@ -124,7 +124,9 @@ def summarize_tool_result(event: ChatEvent) -> str:
124
  if not matches:
125
  try:
126
  payload = json.loads(str(event.data.get("output_text", "")))
127
- matches = payload.get("matches", []) if isinstance(payload, dict) else []
 
 
128
  match_count = len(matches)
129
  except json.JSONDecodeError:
130
  matches = []
@@ -221,11 +223,7 @@ def render_activity_block(events: list[ActivityEvent]) -> str:
221
  return ""
222
 
223
  rendered_sections = "\n\n".join(section for section in sections if section)
224
- return (
225
- f"{ACTIVITY_BLOCK_START}\n\n"
226
- f"{rendered_sections}\n\n"
227
- f"{ACTIVITY_BLOCK_END}"
228
- )
229
 
230
 
231
  def format_sources(matches_by_doc_id: dict[str, dict[str, Any]]) -> str:
@@ -273,7 +271,9 @@ class GradioPresenterState:
273
  activity_events: list[ActivityEvent] = field(default_factory=list)
274
  answer_chunks: list[str] = field(default_factory=list)
275
  completed_answer: str = ""
276
- tool_matches_by_call_id: dict[str, list[dict[str, Any]]] = field(default_factory=dict)
 
 
277
 
278
  def apply(self, event: ChatEvent) -> None:
279
  if event.type == "thread_started":
 
124
  if not matches:
125
  try:
126
  payload = json.loads(str(event.data.get("output_text", "")))
127
+ matches = (
128
+ payload.get("matches", []) if isinstance(payload, dict) else []
129
+ )
130
  match_count = len(matches)
131
  except json.JSONDecodeError:
132
  matches = []
 
223
  return ""
224
 
225
  rendered_sections = "\n\n".join(section for section in sections if section)
226
+ return f"{ACTIVITY_BLOCK_START}\n\n{rendered_sections}\n\n{ACTIVITY_BLOCK_END}"
 
 
 
 
227
 
228
 
229
  def format_sources(matches_by_doc_id: dict[str, dict[str, Any]]) -> str:
 
271
  activity_events: list[ActivityEvent] = field(default_factory=list)
272
  answer_chunks: list[str] = field(default_factory=list)
273
  completed_answer: str = ""
274
+ tool_matches_by_call_id: dict[str, list[dict[str, Any]]] = field(
275
+ default_factory=dict
276
+ )
277
 
278
  def apply(self, event: ChatEvent) -> None:
279
  if event.type == "thread_started":
scripts/kb_manifest.py CHANGED
@@ -62,7 +62,12 @@ def load_manifest_entries(kb_dir: str = KB_DIR) -> tuple[KbManifestEntry, ...]:
62
 
63
  def manifest_indexes(
64
  kb_dir: str = KB_DIR,
65
- ) -> tuple[dict[str, KbManifestEntry], dict[str, KbManifestEntry], dict[str, KbManifestEntry], dict[str, KbManifestEntry]]:
 
 
 
 
 
66
  by_doc_id: dict[str, KbManifestEntry] = {}
67
  by_url: dict[str, KbManifestEntry] = {}
68
  by_path: dict[str, KbManifestEntry] = {}
@@ -81,7 +86,9 @@ def manifest_indexes(
81
  return by_doc_id, by_url, by_path, by_title
82
 
83
 
84
- def source_match_from_manifest(entry: KbManifestEntry, *, score: float = 1.0) -> SourceMatch:
 
 
85
  return SourceMatch(
86
  doc_id=entry.doc_id,
87
  title=entry.title,
@@ -89,7 +96,9 @@ def source_match_from_manifest(entry: KbManifestEntry, *, score: float = 1.0) ->
89
  source_key=entry.source,
90
  source_label=SOURCE_KEY_TO_LABEL.get(entry.source, entry.source),
91
  score=score,
92
- group="courses" if entry.source in COURSE_SOURCE_KEYS else entry.source_group or "docs",
 
 
93
  )
94
 
95
 
@@ -167,7 +176,9 @@ def citation_dedupe_key(match: SourceMatch) -> str:
167
  return normalize_url(match.url) or match.doc_id or match.title
168
 
169
 
170
- def source_match_payload(match: SourceMatch, *, message_id: str, call_id: str = "") -> dict[str, Any]:
 
 
171
  payload: dict[str, Any] = {
172
  "message_id": message_id,
173
  "doc_id": match.doc_id,
 
62
 
63
  def manifest_indexes(
64
  kb_dir: str = KB_DIR,
65
+ ) -> tuple[
66
+ dict[str, KbManifestEntry],
67
+ dict[str, KbManifestEntry],
68
+ dict[str, KbManifestEntry],
69
+ dict[str, KbManifestEntry],
70
+ ]:
71
  by_doc_id: dict[str, KbManifestEntry] = {}
72
  by_url: dict[str, KbManifestEntry] = {}
73
  by_path: dict[str, KbManifestEntry] = {}
 
86
  return by_doc_id, by_url, by_path, by_title
87
 
88
 
89
+ def source_match_from_manifest(
90
+ entry: KbManifestEntry, *, score: float = 1.0
91
+ ) -> SourceMatch:
92
  return SourceMatch(
93
  doc_id=entry.doc_id,
94
  title=entry.title,
 
96
  source_key=entry.source,
97
  source_label=SOURCE_KEY_TO_LABEL.get(entry.source, entry.source),
98
  score=score,
99
+ group="courses"
100
+ if entry.source in COURSE_SOURCE_KEYS
101
+ else entry.source_group or "docs",
102
  )
103
 
104
 
 
176
  return normalize_url(match.url) or match.doc_id or match.title
177
 
178
 
179
+ def source_match_payload(
180
+ match: SourceMatch, *, message_id: str, call_id: str = ""
181
+ ) -> dict[str, Any]:
182
  payload: dict[str, Any] = {
183
  "message_id": message_id,
184
  "doc_id": match.doc_id,
scripts/kb_shell.py CHANGED
@@ -41,7 +41,9 @@ RG_FLAG_OPTIONS = frozenset(
41
  "--no-ignore",
42
  }
43
  )
44
- GREP_FLAG_OPTIONS = frozenset({"-i", "--ignore-case", "-w", "--word-regexp", "-F", "-E"})
 
 
45
  LS_FLAG_OPTIONS = frozenset({"-1", "-a", "-l", "-la", "-al"})
46
  WC_FLAG_OPTIONS = frozenset({"-l", "-w", "-c", "-m"})
47
 
@@ -117,7 +119,9 @@ def _is_broad_raw_path(path: str) -> bool:
117
  return normalized in {".", "raw", "raw/courses", "raw/docs"}
118
 
119
 
120
- def _reject_unbounded_raw_search(command: str, paths: list[str], has_max_count: bool) -> None:
 
 
121
  if has_max_count:
122
  return
123
  if any(_is_broad_raw_path(path) for path in paths):
@@ -230,7 +234,9 @@ def _build_find(tokens: list[str], root: Path) -> list[str]:
230
  argv.extend([token, _safe_pattern_value(tokens[idx + 1], token)])
231
  idx += 2
232
  continue
233
- raise KbCommandError("`find` supports only path, -maxdepth, -mindepth, -type, -name, and -iname")
 
 
234
  return argv
235
 
236
 
@@ -291,7 +297,9 @@ def _build_wc(tokens: list[str], root: Path) -> list[str]:
291
  return [*argv, *paths]
292
 
293
 
294
- def build_kb_command_argv(command: str, *, root: Path | None = None) -> tuple[list[str], Path]:
 
 
295
  resolved_root = _resolve_root(root)
296
  try:
297
  tokens = shlex.split(command)
@@ -304,7 +312,9 @@ def build_kb_command_argv(command: str, *, root: Path | None = None) -> tuple[li
304
  executable = tokens[0]
305
  if executable not in SUPPORTED_COMMANDS:
306
  supported = ", ".join(sorted(SUPPORTED_COMMANDS))
307
- raise KbCommandError(f"Unsupported KB command `{executable}`. Supported: {supported}")
 
 
308
 
309
  builders = {
310
  "rg": _build_rg,
@@ -361,7 +371,9 @@ def run_kb_command(
361
  exit_code = 124
362
  stdout = exc.stdout if isinstance(exc.stdout, str) else ""
363
  stderr = exc.stderr if isinstance(exc.stderr, str) else ""
364
- stderr = (stderr + "\n" if stderr else "") + f"Command timed out after {timeout}s."
 
 
365
 
366
  stdout, stdout_truncated = _truncate(stdout, output_limit)
367
  stderr, stderr_truncated = _truncate(stderr, min(output_limit, 8_000))
 
41
  "--no-ignore",
42
  }
43
  )
44
+ GREP_FLAG_OPTIONS = frozenset(
45
+ {"-i", "--ignore-case", "-w", "--word-regexp", "-F", "-E"}
46
+ )
47
  LS_FLAG_OPTIONS = frozenset({"-1", "-a", "-l", "-la", "-al"})
48
  WC_FLAG_OPTIONS = frozenset({"-l", "-w", "-c", "-m"})
49
 
 
119
  return normalized in {".", "raw", "raw/courses", "raw/docs"}
120
 
121
 
122
+ def _reject_unbounded_raw_search(
123
+ command: str, paths: list[str], has_max_count: bool
124
+ ) -> None:
125
  if has_max_count:
126
  return
127
  if any(_is_broad_raw_path(path) for path in paths):
 
234
  argv.extend([token, _safe_pattern_value(tokens[idx + 1], token)])
235
  idx += 2
236
  continue
237
+ raise KbCommandError(
238
+ "`find` supports only path, -maxdepth, -mindepth, -type, -name, and -iname"
239
+ )
240
  return argv
241
 
242
 
 
297
  return [*argv, *paths]
298
 
299
 
300
+ def build_kb_command_argv(
301
+ command: str, *, root: Path | None = None
302
+ ) -> tuple[list[str], Path]:
303
  resolved_root = _resolve_root(root)
304
  try:
305
  tokens = shlex.split(command)
 
312
  executable = tokens[0]
313
  if executable not in SUPPORTED_COMMANDS:
314
  supported = ", ".join(sorted(SUPPORTED_COMMANDS))
315
+ raise KbCommandError(
316
+ f"Unsupported KB command `{executable}`. Supported: {supported}"
317
+ )
318
 
319
  builders = {
320
  "rg": _build_rg,
 
371
  exit_code = 124
372
  stdout = exc.stdout if isinstance(exc.stdout, str) else ""
373
  stderr = exc.stderr if isinstance(exc.stderr, str) else ""
374
+ stderr = (
375
+ stderr + "\n" if stderr else ""
376
+ ) + f"Command timed out after {timeout}s."
377
 
378
  stdout, stdout_truncated = _truncate(stdout, output_limit)
379
  stderr, stderr_truncated = _truncate(stderr, min(output_limit, 8_000))
scripts/probe_gemini_cache.py CHANGED
@@ -69,12 +69,16 @@ def main() -> None:
69
  client = genai.Client(api_key=api_key)
70
 
71
  print(f"Model: {MODEL}")
72
- print(f"Stable prefix (system + LONG_CONTEXT) repeats. Only the final user_text differs.")
 
 
73
  print()
74
 
75
  for i in range(3):
76
- result = call(client, user_text=f"Question {i+1}: what is RAG in one sentence?")
77
- print(f"Call {i+1}: {result}")
 
 
78
  time.sleep(0.5)
79
 
80
 
 
69
  client = genai.Client(api_key=api_key)
70
 
71
  print(f"Model: {MODEL}")
72
+ print(
73
+ "Stable prefix (system + LONG_CONTEXT) repeats. Only the final user_text differs."
74
+ )
75
  print()
76
 
77
  for i in range(3):
78
+ result = call(
79
+ client, user_text=f"Question {i + 1}: what is RAG in one sentence?"
80
+ )
81
+ print(f"Call {i + 1}: {result}")
82
  time.sleep(0.5)
83
 
84
 
scripts/prompts.py CHANGED
@@ -205,7 +205,9 @@ def build_system_prompt(model_name: str, enabled_tools: tuple[str, ...]) -> str:
205
  KB_USAGE_SECTION,
206
  ]
207
  if usage_sections:
208
- parts.append("## When to use web search / URL reading\n\n" + "\n\n".join(usage_sections))
 
 
209
  parts.append(
210
  "Prefer `retrieve_tutor_context` first when the question is clearly about\n"
211
  "course material. Combine tools when it helps (e.g. retrieve corpus\n"
 
205
  KB_USAGE_SECTION,
206
  ]
207
  if usage_sections:
208
+ parts.append(
209
+ "## When to use web search / URL reading\n\n" + "\n\n".join(usage_sections)
210
+ )
211
  parts.append(
212
  "Prefer `retrieve_tutor_context` first when the question is clearly about\n"
213
  "course material. Combine tools when it helps (e.g. retrieve corpus\n"
tests/test_api.py CHANGED
@@ -56,7 +56,10 @@ class ApiTestCase(unittest.TestCase):
56
  any(source["selectedByDefault"] for source in retrieval["sources"])
57
  )
58
  self.assertTrue(
59
- all(source["group"] in {"courses", "docs"} for source in retrieval["sources"])
 
 
 
60
  )
61
  # Gemini is the default model, so web search + url reading are present.
62
  tool_keys = {tool["key"] for tool in tools}
@@ -124,8 +127,13 @@ class ApiTestCase(unittest.TestCase):
124
  "output_text": "payload",
125
  },
126
  )
127
- yield ChatEvent("text_delta", {"message_id": "message_1", "text": "RAG combines "})
128
- yield ChatEvent("text_delta", {"message_id": "message_1", "text": "retrieval with generation."})
 
 
 
 
 
129
  yield ChatEvent(
130
  "message_completed",
131
  {
@@ -189,7 +197,9 @@ class ApiTestCase(unittest.TestCase):
189
 
190
  with patch("scripts.api.stream_chat", broken_stream_chat):
191
  with TestClient(app) as client:
192
- with client.stream("POST", "/api/chat", json={"query": "Hello"}) as response:
 
 
193
  body = "".join(response.iter_text())
194
 
195
  self.assertEqual(response.status_code, 200)
@@ -208,7 +218,9 @@ class ApiTestCase(unittest.TestCase):
208
  def test_chat_stream_restarts_reasoning_after_tool_activity(self) -> None:
209
  async def fake_stream_chat(_request):
210
  yield ChatEvent("message_started", {"message_id": "message_1"})
211
- yield ChatEvent("reasoning_delta", {"message_id": "message_1", "text": "First thought"})
 
 
212
  yield ChatEvent(
213
  "tool_call_started",
214
  {
@@ -227,8 +239,12 @@ class ApiTestCase(unittest.TestCase):
227
  "output_text": "payload",
228
  },
229
  )
230
- yield ChatEvent("reasoning_delta", {"message_id": "message_1", "text": "Second thought"})
231
- yield ChatEvent("text_delta", {"message_id": "message_1", "text": "Final answer"})
 
 
 
 
232
  yield ChatEvent(
233
  "message_completed",
234
  {
@@ -239,7 +255,9 @@ class ApiTestCase(unittest.TestCase):
239
 
240
  with patch("scripts.api.stream_chat", fake_stream_chat):
241
  with TestClient(app) as client:
242
- with client.stream("POST", "/api/chat", json={"query": "Hello"}) as response:
 
 
243
  body = "".join(response.iter_text())
244
 
245
  self.assertEqual(response.status_code, 200)
@@ -249,7 +267,9 @@ class ApiTestCase(unittest.TestCase):
249
 
250
  self.assertEqual(part_types.count("reasoning-start"), 2)
251
  self.assertEqual(part_types.count("reasoning-end"), 2)
252
- self.assertLess(part_types.index("tool-input-start"), part_types.index("text-start"))
 
 
253
 
254
 
255
  LIVE_API_E2E = pytest.mark.skipif(
 
56
  any(source["selectedByDefault"] for source in retrieval["sources"])
57
  )
58
  self.assertTrue(
59
+ all(
60
+ source["group"] in {"courses", "docs"}
61
+ for source in retrieval["sources"]
62
+ )
63
  )
64
  # Gemini is the default model, so web search + url reading are present.
65
  tool_keys = {tool["key"] for tool in tools}
 
127
  "output_text": "payload",
128
  },
129
  )
130
+ yield ChatEvent(
131
+ "text_delta", {"message_id": "message_1", "text": "RAG combines "}
132
+ )
133
+ yield ChatEvent(
134
+ "text_delta",
135
+ {"message_id": "message_1", "text": "retrieval with generation."},
136
+ )
137
  yield ChatEvent(
138
  "message_completed",
139
  {
 
197
 
198
  with patch("scripts.api.stream_chat", broken_stream_chat):
199
  with TestClient(app) as client:
200
+ with client.stream(
201
+ "POST", "/api/chat", json={"query": "Hello"}
202
+ ) as response:
203
  body = "".join(response.iter_text())
204
 
205
  self.assertEqual(response.status_code, 200)
 
218
  def test_chat_stream_restarts_reasoning_after_tool_activity(self) -> None:
219
  async def fake_stream_chat(_request):
220
  yield ChatEvent("message_started", {"message_id": "message_1"})
221
+ yield ChatEvent(
222
+ "reasoning_delta", {"message_id": "message_1", "text": "First thought"}
223
+ )
224
  yield ChatEvent(
225
  "tool_call_started",
226
  {
 
239
  "output_text": "payload",
240
  },
241
  )
242
+ yield ChatEvent(
243
+ "reasoning_delta", {"message_id": "message_1", "text": "Second thought"}
244
+ )
245
+ yield ChatEvent(
246
+ "text_delta", {"message_id": "message_1", "text": "Final answer"}
247
+ )
248
  yield ChatEvent(
249
  "message_completed",
250
  {
 
255
 
256
  with patch("scripts.api.stream_chat", fake_stream_chat):
257
  with TestClient(app) as client:
258
+ with client.stream(
259
+ "POST", "/api/chat", json={"query": "Hello"}
260
+ ) as response:
261
  body = "".join(response.iter_text())
262
 
263
  self.assertEqual(response.status_code, 200)
 
267
 
268
  self.assertEqual(part_types.count("reasoning-start"), 2)
269
  self.assertEqual(part_types.count("reasoning-end"), 2)
270
+ self.assertLess(
271
+ part_types.index("tool-input-start"), part_types.index("text-start")
272
+ )
273
 
274
 
275
  LIVE_API_E2E = pytest.mark.skipif(
tests/test_build_kb_artifacts.py CHANGED
@@ -89,7 +89,10 @@ Use `LoraConfig` with `get_peft_model`.
89
  assert manifest_rows[0]["original_path"] == docs_page.as_posix()
90
 
91
  markdown_path = Path(manifest_rows[0]["path"])
92
- assert markdown_path == output_dir / "raw" / "docs" / "temp_docs" / "package_reference" / "lora.mdx"
 
 
 
93
  markdown = markdown_path.read_text(encoding="utf-8")
94
  assert 'doc_id: "temp_docs:docs-package-reference-lora"' in markdown
95
  assert 'source_path: "package_reference/lora.mdx"' in markdown
 
89
  assert manifest_rows[0]["original_path"] == docs_page.as_posix()
90
 
91
  markdown_path = Path(manifest_rows[0]["path"])
92
+ assert (
93
+ markdown_path
94
+ == output_dir / "raw" / "docs" / "temp_docs" / "package_reference" / "lora.mdx"
95
+ )
96
  markdown = markdown_path.read_text(encoding="utf-8")
97
  assert 'doc_id: "temp_docs:docs-package-reference-lora"' in markdown
98
  assert 'source_path: "package_reference/lora.mdx"' in markdown
tests/test_chat_service.py CHANGED
@@ -209,8 +209,12 @@ class ChatServiceTestCase(unittest.TestCase):
209
  )
210
  # Enabling Gemini web tools adds exactly one toggle-specific middleware
211
  # (GeminiServerSideToolsMiddleware) on top of the shared base middlewares.
212
- web_middleware = [type(m).__name__ for m in created_agents[0].kwargs["middleware"]]
213
- plain_middleware = [type(m).__name__ for m in created_agents[1].kwargs["middleware"]]
 
 
 
 
214
  self.assertIn("GeminiServerSideToolsMiddleware", web_middleware)
215
  self.assertNotIn("GeminiServerSideToolsMiddleware", plain_middleware)
216
  self.assertEqual(len(web_middleware), len(plain_middleware) + 1)
@@ -382,11 +386,7 @@ class ChatServiceTestCase(unittest.TestCase):
382
  if event.type == "tool_call_completed"
383
  and event.data.get("tool_name") == "run_kb_command"
384
  ]
385
- source_matches = [
386
- event
387
- for event in events
388
- if event.type == "source_match"
389
- ]
390
 
391
  self.assertEqual(started[0].data["args_text"], "rg LoraConfig raw")
392
  self.assertIn("rg LoraConfig raw", completed[0].data["output_text"])
 
209
  )
210
  # Enabling Gemini web tools adds exactly one toggle-specific middleware
211
  # (GeminiServerSideToolsMiddleware) on top of the shared base middlewares.
212
+ web_middleware = [
213
+ type(m).__name__ for m in created_agents[0].kwargs["middleware"]
214
+ ]
215
+ plain_middleware = [
216
+ type(m).__name__ for m in created_agents[1].kwargs["middleware"]
217
+ ]
218
  self.assertIn("GeminiServerSideToolsMiddleware", web_middleware)
219
  self.assertNotIn("GeminiServerSideToolsMiddleware", plain_middleware)
220
  self.assertEqual(len(web_middleware), len(plain_middleware) + 1)
 
386
  if event.type == "tool_call_completed"
387
  and event.data.get("tool_name") == "run_kb_command"
388
  ]
389
+ source_matches = [event for event in events if event.type == "source_match"]
 
 
 
 
390
 
391
  self.assertEqual(started[0].data["args_text"], "rg LoraConfig raw")
392
  self.assertIn("rg LoraConfig raw", completed[0].data["output_text"])
tests/test_chroma_rag.py CHANGED
@@ -86,7 +86,9 @@ After the example.
86
  ]
87
  index = BM25Index.build(records)
88
 
89
- hits = index.search("AutoModel.from_pretrained", allowed_sources=["transformers"])
 
 
90
 
91
  self.assertEqual([record.chunk_id for record, _score in hits], ["a"])
92
  self.assertEqual(index.search("AutoModel", allowed_sources=["langchain"]), [])
@@ -159,7 +161,9 @@ After the example.
159
 
160
  self.assertGreaterEqual(len(records), 1)
161
  setup_record = next(
162
- record for record in records if record.metadata["heading_path"] == "Guide > Setup"
 
 
163
  )
164
  self.assertIn("raw_text", setup_record.metadata)
165
  self.assertIn("Title: Guide", setup_record.text)
 
86
  ]
87
  index = BM25Index.build(records)
88
 
89
+ hits = index.search(
90
+ "AutoModel.from_pretrained", allowed_sources=["transformers"]
91
+ )
92
 
93
  self.assertEqual([record.chunk_id for record, _score in hits], ["a"])
94
  self.assertEqual(index.search("AutoModel", allowed_sources=["langchain"]), [])
 
161
 
162
  self.assertGreaterEqual(len(records), 1)
163
  setup_record = next(
164
+ record
165
+ for record in records
166
+ if record.metadata["heading_path"] == "Guide > Setup"
167
  )
168
  self.assertIn("raw_text", setup_record.metadata)
169
  self.assertIn("Title: Guide", setup_record.text)
tests/test_gradio_kb_e2e.py CHANGED
@@ -172,7 +172,9 @@ def test_live_gradio_curl_can_use_retrieval_and_shell() -> None:
172
  "Use both retrieve_tutor_context and run_kb_command to answer: how does PEFT configure LoRA with LoraConfig?",
173
  [],
174
  ["PEFT Docs", "Transformers Docs"],
175
- os.getenv("LIVE_GRADIO_E2E_MODEL", "google-genai:gemini-3.5-flash"),
 
 
176
  "",
177
  False,
178
  False,
 
172
  "Use both retrieve_tutor_context and run_kb_command to answer: how does PEFT configure LoRA with LoraConfig?",
173
  [],
174
  ["PEFT Docs", "Transformers Docs"],
175
+ os.getenv(
176
+ "LIVE_GRADIO_E2E_MODEL", "google-genai:gemini-3.5-flash"
177
+ ),
178
  "",
179
  False,
180
  False,
tests/test_gradio_presenter.py CHANGED
@@ -56,8 +56,7 @@ class GradioPresenterTestCase(unittest.TestCase):
56
  {
57
  "tool_name": "retrieve_tutor_context",
58
  "output_text": (
59
- '{"matches": [{"title": "LoRA", '
60
- '"source_label": "PEFT Docs"}]}'
61
  ),
62
  },
63
  )
 
56
  {
57
  "tool_name": "retrieve_tutor_context",
58
  "output_text": (
59
+ '{"matches": [{"title": "LoRA", "source_label": "PEFT Docs"}]}'
 
60
  ),
61
  },
62
  )
tests/test_retire_source_workflow.py CHANGED
@@ -6,7 +6,7 @@ from data.scraping_scripts.retire_source_workflow import (
6
 
7
 
8
  def test_remove_source_from_registry_text_removes_active_source_entries():
9
- registry = '''SOURCE_CONFIGS = {
10
  "openai_cookbooks": {
11
  "output_file": "data/openai_cookbooks_data.jsonl",
12
  "nested": {"keep": "balanced"},
@@ -32,7 +32,7 @@ UI_SOURCE_KEYS = (
32
  "openai_cookbooks",
33
  "langchain",
34
  )
35
- '''
36
 
37
  updated, changed = remove_source_from_registry_text(
38
  registry,
 
6
 
7
 
8
  def test_remove_source_from_registry_text_removes_active_source_entries():
9
+ registry = """SOURCE_CONFIGS = {
10
  "openai_cookbooks": {
11
  "output_file": "data/openai_cookbooks_data.jsonl",
12
  "nested": {"keep": "balanced"},
 
32
  "openai_cookbooks",
33
  "langchain",
34
  )
35
+ """
36
 
37
  updated, changed = remove_source_from_registry_text(
38
  registry,
tests/test_update_kb_wiki.py CHANGED
@@ -37,7 +37,7 @@ def test_update_kb_wiki_seeds_navigation_pages(tmp_path: Path) -> None:
37
  "tokens": 300,
38
  "retrieve_doc": True,
39
  "content": "# Lesson 1\n\nAgent lesson.",
40
- }
41
  ],
42
  )
43
  build_kb_artifacts(input_file, kb_dir)
@@ -61,7 +61,9 @@ def test_update_kb_wiki_seeds_navigation_pages(tmp_path: Path) -> None:
61
  assert "lookup_tutor_symbol" not in topic
62
 
63
 
64
- def test_update_kb_wiki_preserves_authored_topic_content_and_appends_log(tmp_path: Path) -> None:
 
 
65
  input_file = tmp_path / "all_sources_data.jsonl"
66
  kb_dir = tmp_path / "kb"
67
  write_jsonl(
 
37
  "tokens": 300,
38
  "retrieve_doc": True,
39
  "content": "# Lesson 1\n\nAgent lesson.",
40
+ },
41
  ],
42
  )
43
  build_kb_artifacts(input_file, kb_dir)
 
61
  assert "lookup_tutor_symbol" not in topic
62
 
63
 
64
+ def test_update_kb_wiki_preserves_authored_topic_content_and_appends_log(
65
+ tmp_path: Path,
66
+ ) -> None:
67
  input_file = tmp_path / "all_sources_data.jsonl"
68
  kb_dir = tmp_path / "kb"
69
  write_jsonl(