ruff format
Browse files- data/scraping_scripts/add_context_to_nodes.py +8 -8
- data/scraping_scripts/build_kb_artifacts.py +21 -10
- data/scraping_scripts/create_vector_stores.py +3 -1
- data/scraping_scripts/github_to_markdown_ai_docs.py +2 -4
- data/scraping_scripts/retire_source_workflow.py +14 -7
- data/scraping_scripts/update_kb_wiki.py +17 -11
- scripts/api.py +17 -8
- scripts/chat_service.py +11 -17
- scripts/chroma_rag.py +13 -7
- scripts/gradio_presenter.py +7 -7
- scripts/kb_manifest.py +15 -4
- scripts/kb_shell.py +18 -6
- scripts/probe_gemini_cache.py +7 -3
- scripts/prompts.py +3 -1
- tests/test_api.py +29 -9
- tests/test_build_kb_artifacts.py +4 -1
- tests/test_chat_service.py +7 -7
- tests/test_chroma_rag.py +6 -2
- tests/test_gradio_kb_e2e.py +3 -1
- tests/test_gradio_presenter.py +1 -2
- tests/test_retire_source_workflow.py +2 -2
- tests/test_update_kb_wiki.py +4 -2
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
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 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
|
| 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 =
|
| 377 |
-
|
|
|
|
|
|
|
|
|
|
| 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] =
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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 = [
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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 |
-
|
| 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(
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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 =
|
|
|
|
|
|
|
| 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 =
|
| 142 |
-
|
| 143 |
-
|
|
|
|
|
|
|
| 144 |
)
|
| 145 |
-
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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 =
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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
|
| 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
|
| 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
|
|
|
|
|
|
|
|
|
|
| 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
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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 =
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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[
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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"
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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 = (
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 73 |
print()
|
| 74 |
|
| 75 |
for i in range(3):
|
| 76 |
-
result = call(
|
| 77 |
-
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
|
|
|
| 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(
|
| 128 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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(
|
| 231 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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
|
|
|
|
|
|
|
|
|
|
| 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 = [
|
| 213 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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 =
|
| 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(
|
|
|
|
|
|
|
| 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(
|