omarsol commited on
Commit
7478ade
·
1 Parent(s): 62d6725

Refactor scraping scripts to enhance source management and contextual node processing. Introduce new helper functions for pruning inactive contextual nodes and managing llms.txt sources. Update existing scripts to utilize the new source registry for improved source handling and streamline the process of adding and updating courses. Enhance error handling and logging for better traceability during data processing.

Browse files
data/scraping_scripts/README.md CHANGED
@@ -4,11 +4,18 @@
4
 
5
  Make sure you have the required environment variables set:
6
 
7
- - `OPENAI_API_KEY` for the LLM
8
  - `COHERE_API_KEY` for embeddings
9
  - `HF_TOKEN` for HuggingFace uploads and downloads - [access to the private HuggingFace dataset repo](https://huggingface.co/datasets/towardsai-tutors/ai-tutor-data/tree/main)
10
  - `GITHUB_TOKEN` for accessing files via the GitHub API
11
 
 
 
 
 
 
 
 
12
  ## 1. Prepare the course data
13
 
14
  0. Make sure you have access to:
@@ -40,9 +47,10 @@ Make sure you have the required environment variables set:
40
  7. Rename the folder to the course name.
41
  - e.g. `master_ai_for_work`
42
 
43
- 8. Open the `data/scraping_scripts/process_md_files.py` python file and locate the `SOURCE_CONFIGS` dictionary.
44
 
45
- 9. Add the new course to the `SOURCE_CONFIGS` dictionary.
 
46
 
47
  example:
48
 
@@ -70,13 +78,13 @@ Make sure you have the required environment variables set:
70
  ## 2. Run the add_course_workflow.py script
71
 
72
  ```bash
73
- uv run -m data.scraping_scripts.add_course_workflow --course [COURSE_NAME]
74
  ```
75
 
76
  example:
77
 
78
  ```bash
79
- uv run -m data.scraping_scripts.add_course_workflow --course master_ai_for_work
80
  ```
81
 
82
  This script will guide you through the complete process, it will:
@@ -88,7 +96,7 @@ This script will guide you through the complete process, it will:
88
  5. Add contextual information to document chunks before embedding
89
  6. Create vector stores
90
  7. Upload databases to HuggingFace
91
- 8. Update UI configuration
92
 
93
  ## 3. Add URLs to the course content + Manual Dataset Cleaning (Most important step)
94
 
@@ -109,7 +117,7 @@ example: "Course Admin and Syllabus", "Course Structure/Overview", "Course Outli
109
  ## 4. Once done, run the script again and answer "yes" to the question "Have you added all the URLs?"
110
 
111
  ```bash
112
- uv run -m data.scraping_scripts.add_course_workflow --course master_ai_for_work
113
  ```
114
 
115
  ## 5. Its done
@@ -172,9 +180,11 @@ The workflow:
172
  1. Downloads `all_sources_data.jsonl` and `all_sources_contextual_nodes.pkl` if
173
  they are missing locally.
174
  2. Removes rows/chunks whose `source` matches the retired source key.
175
- 3. Rebuilds `data/chroma-db-all_sources`.
176
- 4. Uploads the rebuilt vector DB and updated aggregate data files.
177
- 5. Deletes the retired per-source JSONL from `towardsai-tutors/ai-tutor-data`.
 
 
178
 
179
  Run with `--dry-run` first to preview counts without changing files:
180
 
@@ -182,8 +192,8 @@ Run with `--dry-run` first to preview counts without changing files:
182
  uv run -m data.scraping_scripts.retire_source_workflow --sources 8-hour_primer --dry-run
183
  ```
184
 
185
- After retiring a source, remove its entry from `SOURCE_CONFIGS` and from the
186
- workflow download/upload lists so future merges cannot reintroduce it.
187
 
188
  ## Tips for New Team Members
189
 
@@ -200,9 +210,9 @@ workflow download/upload lists so future merges cannot reintroduce it.
200
  3. By default, only new content will have context added to save time and resources. Use `--process-all-context` only if you need to regenerate context for all documents. Use `--skip-data-upload` if you don't want to upload data files to the private HuggingFace repo (they're uploaded by default).
201
 
202
  4. When adding a new course, verify that it appears in the Gradio UI:
203
- - The workflow automatically updates `scripts/setup.py` to include the new source and `scripts/main.py` to preselect it by default
204
  - Check that the new source appears in the dropdown menu in the UI
205
- - Make sure it's properly included in the default selected sources
206
  - Restart the Gradio app to see the changes
207
 
208
  5. First time setup or missing files:
 
4
 
5
  Make sure you have the required environment variables set:
6
 
7
+ - `GEMINI_API_KEY` or `GOOGLE_API_KEY` for context generation with Gemini
8
  - `COHERE_API_KEY` for embeddings
9
  - `HF_TOKEN` for HuggingFace uploads and downloads - [access to the private HuggingFace dataset repo](https://huggingface.co/datasets/towardsai-tutors/ai-tutor-data/tree/main)
10
  - `GITHUB_TOKEN` for accessing files via the GitHub API
11
 
12
+ Optional Gemini context-generation tuning:
13
+
14
+ - `GEMINI_CONTEXT_TPM_LIMIT` - input token-per-minute quota to throttle against (defaults to `30000000`)
15
+ - `GEMINI_CONTEXT_TPM_SAFETY_MARGIN` - fraction of that quota to use before pausing (defaults to `0.8`)
16
+ - `GEMINI_CONTEXT_CONCURRENCY` - max concurrent context requests (defaults to `50`)
17
+ - `GEMINI_CONTEXT_RETRY_ATTEMPTS` - max tenacity attempts for transient Gemini API errors (defaults to `8`)
18
+
19
  ## 1. Prepare the course data
20
 
21
  0. Make sure you have access to:
 
47
  7. Rename the folder to the course name.
48
  - e.g. `master_ai_for_work`
49
 
50
+ 8. Open `data/scraping_scripts/source_registry.py`.
51
 
52
+ 9. Add the new course to the `SOURCE_CONFIGS` dictionary. Sources listed in
53
+ this registry are active in the knowledge base.
54
 
55
  example:
56
 
 
78
  ## 2. Run the add_course_workflow.py script
79
 
80
  ```bash
81
+ uv run -m data.scraping_scripts.add_course_workflow --courses [COURSE_NAME]
82
  ```
83
 
84
  example:
85
 
86
  ```bash
87
+ uv run -m data.scraping_scripts.add_course_workflow --courses master_ai_for_work
88
  ```
89
 
90
  This script will guide you through the complete process, it will:
 
96
  5. Add contextual information to document chunks before embedding
97
  6. Create vector stores
98
  7. Upload databases to HuggingFace
99
+ 8. Confirm the course is configured in the central source registry
100
 
101
  ## 3. Add URLs to the course content + Manual Dataset Cleaning (Most important step)
102
 
 
117
  ## 4. Once done, run the script again and answer "yes" to the question "Have you added all the URLs?"
118
 
119
  ```bash
120
+ uv run -m data.scraping_scripts.add_course_workflow --courses master_ai_for_work
121
  ```
122
 
123
  ## 5. Its done
 
180
  1. Downloads `all_sources_data.jsonl` and `all_sources_contextual_nodes.pkl` if
181
  they are missing locally.
182
  2. Removes rows/chunks whose `source` matches the retired source key.
183
+ 3. Removes the retired source from `source_registry.py`, so it is no longer
184
+ active in future workflow runs or the UI source picker.
185
+ 4. Rebuilds `data/chroma-db-all_sources`.
186
+ 5. Uploads the rebuilt vector DB and updated aggregate data files.
187
+ 6. Deletes the retired per-source JSONL from `towardsai-tutors/ai-tutor-data`.
188
 
189
  Run with `--dry-run` first to preview counts without changing files:
190
 
 
192
  uv run -m data.scraping_scripts.retire_source_workflow --sources 8-hour_primer --dry-run
193
  ```
194
 
195
+ Use `--keep-source-registry` only if you want to remove existing chunks while
196
+ leaving the source configured as active for a future rebuild.
197
 
198
  ## Tips for New Team Members
199
 
 
210
  3. By default, only new content will have context added to save time and resources. Use `--process-all-context` only if you need to regenerate context for all documents. Use `--skip-data-upload` if you don't want to upload data files to the private HuggingFace repo (they're uploaded by default).
211
 
212
  4. When adding a new course, verify that it appears in the Gradio UI:
213
+ - Add the source label and default-selection metadata in `source_registry.py`
214
  - Check that the new source appears in the dropdown menu in the UI
215
+ - Make sure it's properly included in the default selected sources if desired
216
  - Restart the Gradio app to see the changes
217
 
218
  5. First time setup or missing files:
data/scraping_scripts/add_context_to_nodes.py CHANGED
@@ -1,26 +1,83 @@
1
  import asyncio
2
  import json
 
3
  import pickle
 
 
 
 
4
  from typing import List
5
 
6
- import instructor
7
  import tiktoken
8
- # from anthropic import AsyncAnthropic
9
  from dotenv import load_dotenv
 
 
 
10
  from jinja2 import Template
11
  from llama_index.core import Document
12
- from llama_index.core.ingestion import IngestionPipeline
13
- from llama_index.core.node_parser import SentenceSplitter
14
- from llama_index.core.schema import TextNode
15
- from openai import AsyncOpenAI
16
  from pydantic import BaseModel, Field
17
- from tenacity import retry, stop_after_attempt, wait_exponential
18
  from tqdm.asyncio import tqdm
19
 
20
- from scripts.chroma_rag import ChunkRecord
21
 
22
  load_dotenv(".env")
23
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
24
  def create_docs(input_file: str) -> List[Document]:
25
  with open(input_file, "r") as f:
26
  documents: list[Document] = []
@@ -61,19 +118,90 @@ class SituatedContext(BaseModel):
61
  )
62
 
63
 
64
- # client = AsyncInstructor(
65
- # client=AsyncAnthropic(),
66
- # create=patch(
67
- # create=AsyncAnthropic().beta.prompt_caching.messages.create,
68
- # mode=Mode.ANTHROPIC_TOOLS,
69
- # ),
70
- # mode=Mode.ANTHROPIC_TOOLS,
71
- # )
72
- aclient = AsyncOpenAI()
73
- client: instructor.AsyncInstructor = instructor.from_openai(aclient)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
74
 
 
75
 
76
- @retry(stop=stop_after_attempt(5), wait=wait_exponential(multiplier=1, min=4, max=10))
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
77
  async def situate_context(doc: str, chunk: str) -> str:
78
  template = Template(
79
  """
@@ -88,85 +216,104 @@ Here is the chunk we want to situate within the whole document above:
88
  </chunk>
89
 
90
  Please give a short succinct context to situate this chunk within the overall document for the purposes of improving search retrieval of the chunk.
91
- Answer only with the succinct context and nothing else.
92
  """
93
  )
94
 
95
  content = template.render(doc=doc, chunk=chunk)
96
-
97
- response = await client.chat.completions.create(
98
- model="gpt-4o-mini",
99
- max_tokens=1000,
100
- temperature=0,
101
- messages=[
102
- {
103
- "role": "user",
104
- "content": content,
105
- }
106
- ],
107
- response_model=SituatedContext,
108
- )
109
- return response.context
110
-
111
-
112
- def node_to_chunk_record(node: TextNode) -> ChunkRecord:
113
- doc_id = str(node.source_node.node_id) # type: ignore[union-attr]
114
- metadata = dict(node.metadata)
115
- metadata["doc_id"] = doc_id
116
- return ChunkRecord(
117
- chunk_id=str(node.node_id),
118
- doc_id=doc_id,
119
- text=str(node.text),
120
- metadata=metadata,
121
  )
122
-
123
-
124
- async def process_chunk(node: TextNode, document_dict: dict[str, Document]) -> ChunkRecord:
125
- doc_id: str = node.source_node.node_id # type: ignore
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
126
  doc: Document = document_dict[doc_id]
127
 
128
- if doc.metadata["tokens"] > 120_000:
129
  # Tokenize the document text
130
- encoding = tiktoken.encoding_for_model("gpt-4o-mini")
131
  tokens = encoding.encode(doc.get_content())
132
 
133
  # Trim to 120,000 tokens
134
- trimmed_tokens = tokens[:120_000]
135
 
136
  # Decode back to text
137
  trimmed_text = encoding.decode(trimmed_tokens)
138
 
139
  # Update the document with trimmed text
140
  doc = Document(text=trimmed_text, metadata=doc.metadata)
141
- doc.metadata["tokens"] = 120_000
142
-
143
- context: str = await situate_context(doc.get_content(), node.text)
144
- node.text = f"{node.text}\n\n{context}"
145
- return node_to_chunk_record(node)
 
 
 
 
 
 
 
 
 
 
146
 
147
 
148
  async def process(
149
- documents: List[Document], semaphore_limit: int = 50
150
  ) -> List[ChunkRecord]:
151
 
152
- # From the document, we create chunks
153
- pipeline = IngestionPipeline(
154
- transformations=[SentenceSplitter(chunk_size=800, chunk_overlap=0)]
155
- )
156
- all_nodes: list[TextNode] = pipeline.run(documents=documents, show_progress=True)
157
- print(f"Number of nodes: {len(all_nodes)}")
158
 
159
  document_dict: dict[str, Document] = {doc.doc_id: doc for doc in documents}
160
 
 
 
 
 
 
 
161
  semaphore = asyncio.Semaphore(semaphore_limit)
162
 
163
- async def process_with_semaphore(node):
164
  async with semaphore:
165
- result = await process_chunk(node, document_dict)
166
  await asyncio.sleep(0.1)
167
  return result
168
 
169
- tasks = [process_with_semaphore(node) for node in all_nodes]
170
 
171
  results: List[ChunkRecord] = await tqdm.gather(*tasks, desc="Processing chunks")
172
 
 
1
  import asyncio
2
  import json
3
+ import os
4
  import pickle
5
+ import random
6
+ import re
7
+ import time
8
+ from collections import deque
9
  from typing import List
10
 
 
11
  import tiktoken
 
12
  from dotenv import load_dotenv
13
+ from google import genai
14
+ from google.genai import types
15
+ from google.genai.errors import APIError
16
  from jinja2 import Template
17
  from llama_index.core import Document
 
 
 
 
18
  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
38
+ RETRYABLE_GENAI_STATUS_CODES = {408, 429, 500, 502, 503, 504}
39
+ _genai_client: genai.Client | None = None
40
+ _token_encoding = tiktoken.get_encoding("cl100k_base")
41
+
42
+
43
+ class AsyncTokenWindowLimiter:
44
+ def __init__(self, tokens_per_window: int, window_seconds: float) -> None:
45
+ self.tokens_per_window = max(1, tokens_per_window)
46
+ self.window_seconds = window_seconds
47
+ self._events: deque[tuple[float, int]] = deque()
48
+ self._used_tokens = 0
49
+ self._lock = asyncio.Lock()
50
+
51
+ async def acquire(self, tokens: int) -> None:
52
+ tokens = min(max(1, tokens), self.tokens_per_window)
53
+
54
+ while True:
55
+ async with self._lock:
56
+ now = time.monotonic()
57
+ self._prune(now)
58
+
59
+ if self._used_tokens + tokens <= self.tokens_per_window:
60
+ self._events.append((now, tokens))
61
+ self._used_tokens += tokens
62
+ return
63
+
64
+ oldest_at, _ = self._events[0]
65
+ delay = max(0.1, self.window_seconds - (now - oldest_at))
66
+
67
+ await asyncio.sleep(delay + random.uniform(0.1, 0.75))
68
+
69
+ def _prune(self, now: float) -> None:
70
+ while self._events and now - self._events[0][0] >= self.window_seconds:
71
+ _, tokens = self._events.popleft()
72
+ self._used_tokens -= tokens
73
+
74
+
75
+ _input_token_limiter = AsyncTokenWindowLimiter(
76
+ tokens_per_window=int(CONTEXT_TPM_LIMIT * CONTEXT_TPM_SAFETY_MARGIN),
77
+ window_seconds=CONTEXT_TPM_WINDOW_SECONDS,
78
+ )
79
+
80
+
81
  def create_docs(input_file: str) -> List[Document]:
82
  with open(input_file, "r") as f:
83
  documents: list[Document] = []
 
118
  )
119
 
120
 
121
+ def get_genai_client() -> genai.Client:
122
+ global _genai_client
123
+
124
+ if _genai_client is None:
125
+ api_key = os.getenv("GEMINI_API_KEY") or os.getenv("GOOGLE_API_KEY")
126
+ if not api_key:
127
+ raise RuntimeError(
128
+ "GEMINI_API_KEY or GOOGLE_API_KEY must be set to add chunk context."
129
+ )
130
+ _genai_client = genai.Client(api_key=api_key)
131
+
132
+ return _genai_client
133
+
134
+
135
+ def is_retryable_genai_error(exc: BaseException) -> bool:
136
+ return (
137
+ isinstance(exc, APIError)
138
+ and getattr(exc, "code", None) in RETRYABLE_GENAI_STATUS_CODES
139
+ )
140
+
141
+
142
+ def count_input_tokens(content: str) -> int:
143
+ return len(_token_encoding.encode(content, disallowed_special=()))
144
+
145
+
146
+ def extract_retry_delay_seconds(exc: BaseException | None) -> float | None:
147
+ if not isinstance(exc, APIError):
148
+ return None
149
+
150
+ retry_delay = _find_retry_delay(getattr(exc, "details", None))
151
+ if retry_delay is not None:
152
+ return retry_delay
153
+
154
+ message = getattr(exc, "message", None) or str(exc)
155
+ match = re.search(r"retry in ([0-9.]+)s", message, flags=re.IGNORECASE)
156
+ if match:
157
+ return float(match.group(1))
158
+
159
+ return None
160
+
161
+
162
+ def _find_retry_delay(value) -> float | None:
163
+ if isinstance(value, dict):
164
+ if value.get("@type") == "type.googleapis.com/google.rpc.RetryInfo":
165
+ return _parse_duration_seconds(value.get("retryDelay"))
166
+ for item in value.values():
167
+ retry_delay = _find_retry_delay(item)
168
+ if retry_delay is not None:
169
+ return retry_delay
170
+ elif isinstance(value, list):
171
+ for item in value:
172
+ retry_delay = _find_retry_delay(item)
173
+ if retry_delay is not None:
174
+ return retry_delay
175
+ return None
176
+
177
+
178
+ def _parse_duration_seconds(value) -> float | None:
179
+ if not isinstance(value, str):
180
+ return None
181
+
182
+ match = re.fullmatch(r"([0-9.]+)s", value)
183
+ if match:
184
+ return float(match.group(1))
185
 
186
+ return None
187
 
188
+
189
+ def wait_for_genai_retry(retry_state) -> float:
190
+ exc = retry_state.outcome.exception() if retry_state.outcome else None
191
+ retry_delay = extract_retry_delay_seconds(exc)
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
+
199
+ @retry(
200
+ retry=retry_if_exception(is_retryable_genai_error),
201
+ stop=stop_after_attempt(CONTEXT_RETRY_ATTEMPTS),
202
+ wait=wait_for_genai_retry,
203
+ reraise=True,
204
+ )
205
  async def situate_context(doc: str, chunk: str) -> str:
206
  template = Template(
207
  """
 
216
  </chunk>
217
 
218
  Please give a short succinct context to situate this chunk within the overall document for the purposes of improving search retrieval of the chunk.
219
+ Return a title for the document and the succinct context.
220
  """
221
  )
222
 
223
  content = template.render(doc=doc, chunk=chunk)
224
+ await _input_token_limiter.acquire(count_input_tokens(content))
225
+
226
+ response = await get_genai_client().aio.models.generate_content(
227
+ model=CONTEXT_MODEL,
228
+ contents=content,
229
+ config=types.GenerateContentConfig(
230
+ max_output_tokens=CONTEXT_MAX_OUTPUT_TOKENS,
231
+ response_mime_type="application/json",
232
+ response_schema=SituatedContext,
233
+ ),
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
234
  )
235
+ if isinstance(response.parsed, SituatedContext):
236
+ return response.parsed.context
237
+ if isinstance(response.parsed, dict):
238
+ return SituatedContext.model_validate(response.parsed).context
239
+ if response.text:
240
+ return SituatedContext.model_validate_json(response.text).context
241
+ raise ValueError("Gemini returned no context response text.")
242
+
243
+
244
+ def document_to_row(document: Document) -> dict:
245
+ return {
246
+ "doc_id": document.doc_id,
247
+ "content": document.get_content(),
248
+ "name": document.metadata["title"],
249
+ "url": document.metadata["url"],
250
+ "source": document.metadata["source"],
251
+ "retrieve_doc": document.metadata["retrieve_doc"],
252
+ "tokens": document.metadata["tokens"],
253
+ }
254
+
255
+
256
+ async def process_chunk(
257
+ chunk_record: ChunkRecord,
258
+ document_dict: dict[str, Document],
259
+ ) -> ChunkRecord:
260
+ doc_id = chunk_record.doc_id
261
  doc: Document = document_dict[doc_id]
262
 
263
+ if doc.metadata["tokens"] > MAX_DOCUMENT_TOKENS:
264
  # Tokenize the document text
265
+ encoding = tiktoken.get_encoding("cl100k_base")
266
  tokens = encoding.encode(doc.get_content())
267
 
268
  # Trim to 120,000 tokens
269
+ trimmed_tokens = tokens[:MAX_DOCUMENT_TOKENS]
270
 
271
  # Decode back to text
272
  trimmed_text = encoding.decode(trimmed_tokens)
273
 
274
  # Update the document with trimmed text
275
  doc = Document(text=trimmed_text, metadata=doc.metadata)
276
+ doc.metadata["tokens"] = MAX_DOCUMENT_TOKENS
277
+
278
+ context = await situate_context(doc.get_content(), chunk_record.text)
279
+ metadata = dict(chunk_record.metadata)
280
+ metadata["raw_text"] = chunk_record.text
281
+ contextual_text = (
282
+ f"{format_chunk_for_retrieval(chunk_record.text, metadata)}"
283
+ f"\n\nContext: {context}"
284
+ )
285
+ return ChunkRecord(
286
+ chunk_id=chunk_record.chunk_id,
287
+ doc_id=chunk_record.doc_id,
288
+ text=contextual_text,
289
+ metadata=metadata,
290
+ )
291
 
292
 
293
  async def process(
294
+ documents: List[Document], semaphore_limit: int = DEFAULT_SEMAPHORE_LIMIT
295
  ) -> List[ChunkRecord]:
296
 
297
+ chunk_records = build_chunk_records([document_to_row(doc) for doc in documents])
298
+ print(f"Number of chunks: {len(chunk_records)}")
 
 
 
 
299
 
300
  document_dict: dict[str, Document] = {doc.doc_id: doc for doc in documents}
301
 
302
+ print(
303
+ "Gemini context rate limits: "
304
+ f"{_input_token_limiter.tokens_per_window:,} input tokens/"
305
+ f"{CONTEXT_TPM_WINDOW_SECONDS:g}s, concurrency={semaphore_limit}"
306
+ )
307
+
308
  semaphore = asyncio.Semaphore(semaphore_limit)
309
 
310
+ async def process_with_semaphore(chunk_record: ChunkRecord):
311
  async with semaphore:
312
+ result = await process_chunk(chunk_record, document_dict)
313
  await asyncio.sleep(0.1)
314
  return result
315
 
316
+ tasks = [process_with_semaphore(chunk_record) for chunk_record in chunk_records]
317
 
318
  results: List[ChunkRecord] = await tqdm.gather(*tasks, desc="Processing chunks")
319
 
data/scraping_scripts/add_course_workflow.py CHANGED
@@ -6,8 +6,8 @@ This script guides you through adding or updating one or more courses in the AI
6
 
7
  1. Process course markdown files to create per-course JSONL data
8
  2. MANDATORY MANUAL STEP: Add URLs to each course JSONL
9
- 3. Rebuild all_sources_data.jsonl from the per-source JSONLs listed in SOURCE_CONFIGS
10
- (this naturally drops any retired sources no longer in SOURCE_CONFIGS)
11
  4. Optionally purge retired sources from the contextual-nodes PKL
12
  5. Add contextual information to document nodes (only new docs by default)
13
  6. Create vector stores
@@ -37,13 +37,21 @@ import os
37
  import pickle
38
  import subprocess
39
  import sys
40
- from pathlib import Path
41
  from typing import Dict, List, Set
42
 
43
  from dotenv import load_dotenv
44
  from huggingface_hub import hf_hub_download
45
 
 
 
 
46
  from data.scraping_scripts.hf_auth import HuggingFaceAuthError, validate_hf_access
 
 
 
 
 
 
47
  from scripts.chroma_rag import get_chunk_record_doc_id
48
 
49
  # Load environment variables from .env file
@@ -69,27 +77,10 @@ def ensure_hf_access() -> None:
69
  sys.exit(1)
70
 
71
 
72
- def ensure_required_files_exist():
73
  """Download required data files from HuggingFace if they don't exist locally."""
74
- # List of files to check and download
75
- required_files = {
76
- # Critical files
77
- "data/all_sources_data.jsonl": "all_sources_data.jsonl",
78
- "data/all_sources_contextual_nodes.pkl": "all_sources_contextual_nodes.pkl",
79
- # Documentation source files
80
- "data/transformers_data.jsonl": "transformers_data.jsonl",
81
- "data/peft_data.jsonl": "peft_data.jsonl",
82
- "data/trl_data.jsonl": "trl_data.jsonl",
83
- "data/llama_index_data.jsonl": "llama_index_data.jsonl",
84
- "data/langchain_data.jsonl": "langchain_data.jsonl",
85
- "data/openai_cookbooks_data.jsonl": "openai_cookbooks_data.jsonl",
86
- # Course files
87
- "data/tai_blog_data.jsonl": "tai_blog_data.jsonl",
88
- "data/master_ai_for_work_data.jsonl": "master_ai_for_work_data.jsonl",
89
- "data/agentic_ai_engineering_data.jsonl": "agentic_ai_engineering_data.jsonl",
90
- "data/full_stack_ai_engineering_data.jsonl": "full_stack_ai_engineering_data.jsonl",
91
- "data/beginner_python_for_ai_engineering_data.jsonl": "beginner_python_for_ai_engineering_data.jsonl",
92
- }
93
 
94
  # Critical files that must be downloaded
95
  critical_files = [
@@ -99,6 +90,14 @@ def ensure_required_files_exist():
99
 
100
  # Check and download each file
101
  for local_path, remote_filename in required_files.items():
 
 
 
 
 
 
 
 
102
  if not os.path.exists(local_path):
103
  logger.info(
104
  f"{remote_filename} not found. Attempting to download from HuggingFace..."
@@ -164,9 +163,6 @@ def process_markdown_files(course_name: str) -> str:
164
 
165
  logger.info(f"Successfully processed markdown files for {course_name}")
166
 
167
- # Determine the output file path from process_md_files.py
168
- from data.scraping_scripts.process_md_files import SOURCE_CONFIGS
169
-
170
  if course_name not in SOURCE_CONFIGS:
171
  logger.error(f"Course {course_name} not found in SOURCE_CONFIGS")
172
  sys.exit(1)
@@ -206,12 +202,11 @@ def manual_url_addition(jsonl_path: str) -> None:
206
 
207
 
208
  def rebuild_all_sources(courses: List[str]) -> None:
209
- """Rebuild all_sources_data.jsonl from the per-source JSONLs in SOURCE_CONFIGS.
210
 
211
  Unlike the previous append-style merge, this drops any source whose entry
212
- has been removed from SOURCE_CONFIGS (e.g. retired/renamed courses) and
213
- reloads every other source from its own JSONL. Safe against stale rows
214
- for sources like llm_developper / python_primer that were renamed.
215
  """
216
  from data.scraping_scripts.process_md_files import combine_all_sources
217
 
@@ -415,85 +410,19 @@ def upload_to_huggingface(upload_jsonl: bool = False) -> None:
415
 
416
 
417
  def update_ui_files(course_name: str) -> None:
418
- """Update setup.py and the default selected UI sources for the new course."""
419
- logger.info(f"Updating UI files with new course: {course_name}")
420
-
421
- # Get the source configuration for display name
422
- from data.scraping_scripts.process_md_files import SOURCE_CONFIGS
423
-
424
- if course_name not in SOURCE_CONFIGS:
425
- logger.error(f"Course {course_name} not found in SOURCE_CONFIGS")
426
  return
427
 
428
- # Get a readable display name for the UI
429
- display_name = course_name.replace("_", " ").title()
430
-
431
- # Update setup.py - add to AVAILABLE_SOURCES, AVAILABLE_SOURCES_UI, and SOURCE_UI_TO_KEY
432
- setup_path = Path("scripts/setup.py")
433
- if setup_path.exists():
434
- setup_content = setup_path.read_text()
435
-
436
- # Check if already added
437
- if f'"{course_name}"' in setup_content:
438
- logger.info(f"Course {course_name} already in setup.py")
439
- else:
440
- # Add to AVAILABLE_SOURCES_UI
441
- ui_list_start = setup_content.find("AVAILABLE_SOURCES_UI = [")
442
- ui_list_end = setup_content.find("]", ui_list_start)
443
- new_ui_content = (
444
- setup_content[:ui_list_end]
445
- + f' "{display_name}",\n'
446
- + setup_content[ui_list_end:]
447
- )
448
-
449
- # Add to AVAILABLE_SOURCES
450
- sources_list_start = new_ui_content.find("AVAILABLE_SOURCES = [")
451
- sources_list_end = new_ui_content.find("]", sources_list_start)
452
- new_content = (
453
- new_ui_content[:sources_list_end]
454
- + f' "{course_name}",\n'
455
- + new_ui_content[sources_list_end:]
456
- )
457
-
458
- mapping_start = new_content.find("SOURCE_UI_TO_KEY = {")
459
- mapping_end = new_content.find("}", mapping_start)
460
- if f'"{display_name}": "{course_name}"' not in new_content:
461
- new_content = (
462
- new_content[:mapping_end]
463
- + f' "{display_name}": "{course_name}",\n'
464
- + new_content[mapping_end:]
465
- )
466
-
467
- # Write updated content
468
- setup_path.write_text(new_content)
469
- logger.info(f"Updated setup.py with {course_name}")
470
- else:
471
- logger.warning(f"setup.py not found at {setup_path}")
472
-
473
- # Update main.py - add to the default selected source list
474
- main_path = Path("scripts/main.py")
475
- if main_path.exists():
476
- main_content = main_path.read_text()
477
-
478
- # Check if already added
479
- if f'"{display_name}"' in main_content:
480
- logger.info(f"Course {course_name} already in main.py")
481
- else:
482
- value_start = main_content.find("value=[")
483
- value_end = main_content.find("]", value_start)
484
- new_main_content = main_content
485
- if value_start != -1 and value_end != -1:
486
- new_main_content = (
487
- main_content[: value_start + 7]
488
- + f' "{display_name}",\n'
489
- + main_content[value_start + 7 :]
490
- )
491
-
492
- # Write updated content
493
- main_path.write_text(new_main_content)
494
- logger.info(f"Updated main.py with {course_name}")
495
- else:
496
- logger.warning(f"main.py not found at {main_path}")
497
 
498
 
499
  def main():
@@ -504,7 +433,7 @@ def main():
504
  "--courses",
505
  nargs="+",
506
  required=True,
507
- help="One or more course names to process (must match SOURCE_CONFIGS)",
508
  )
509
  parser.add_argument(
510
  "--purge-sources",
@@ -553,19 +482,19 @@ def main():
553
  args = parser.parse_args()
554
  courses: List[str] = args.courses
555
 
556
- ensure_hf_access()
557
-
558
- # Ensure required data files exist before proceeding
559
- ensure_required_files_exist()
560
-
561
- from data.scraping_scripts.process_md_files import SOURCE_CONFIGS
562
-
563
  # Validate every course up front so we fail fast
564
  for course_name in courses:
565
  if course_name not in SOURCE_CONFIGS:
566
  logger.error(f"Course {course_name} not found in SOURCE_CONFIGS")
567
  sys.exit(1)
568
 
 
 
 
 
 
 
 
569
  # Per-course: process markdown + manual URL addition
570
  for course_name in courses:
571
  course_jsonl_path = SOURCE_CONFIGS[course_name]["output_file"]
@@ -591,6 +520,8 @@ def main():
591
  if not args.skip_context:
592
  add_context_to_nodes(not args.process_all_context)
593
 
 
 
594
  if not args.skip_vectors:
595
  create_vector_stores()
596
 
 
6
 
7
  1. Process course markdown files to create per-course JSONL data
8
  2. MANDATORY MANUAL STEP: Add URLs to each course JSONL
9
+ 3. Rebuild all_sources_data.jsonl from active sources in source_registry.py
10
+ (this naturally drops any retired sources no longer in the registry)
11
  4. Optionally purge retired sources from the contextual-nodes PKL
12
  5. Add contextual information to document nodes (only new docs by default)
13
  6. Create vector stores
 
37
  import pickle
38
  import subprocess
39
  import sys
 
40
  from typing import Dict, List, Set
41
 
42
  from dotenv import load_dotenv
43
  from huggingface_hub import hf_hub_download
44
 
45
+ from data.scraping_scripts.contextual_node_pruning import (
46
+ prune_contextual_nodes_to_active_sources,
47
+ )
48
  from data.scraping_scripts.hf_auth import HuggingFaceAuthError, validate_hf_access
49
+ from data.scraping_scripts.source_registry import (
50
+ SOURCE_CONFIGS,
51
+ SOURCE_KEY_TO_LABEL,
52
+ required_data_files,
53
+ source_output_files,
54
+ )
55
  from scripts.chroma_rag import get_chunk_record_doc_id
56
 
57
  # Load environment variables from .env file
 
77
  sys.exit(1)
78
 
79
 
80
+ def ensure_required_files_exist(sources_to_regenerate: List[str] | None = None):
81
  """Download required data files from HuggingFace if they don't exist locally."""
82
+ required_files = required_data_files()
83
+ regenerated_source_files = source_output_files(sources_to_regenerate or [])
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
84
 
85
  # Critical files that must be downloaded
86
  critical_files = [
 
90
 
91
  # Check and download each file
92
  for local_path, remote_filename in required_files.items():
93
+ if local_path in regenerated_source_files:
94
+ if not os.path.exists(local_path):
95
+ logger.info(
96
+ "%s will be regenerated for this run; skipping HuggingFace download",
97
+ remote_filename,
98
+ )
99
+ continue
100
+
101
  if not os.path.exists(local_path):
102
  logger.info(
103
  f"{remote_filename} not found. Attempting to download from HuggingFace..."
 
163
 
164
  logger.info(f"Successfully processed markdown files for {course_name}")
165
 
 
 
 
166
  if course_name not in SOURCE_CONFIGS:
167
  logger.error(f"Course {course_name} not found in SOURCE_CONFIGS")
168
  sys.exit(1)
 
202
 
203
 
204
  def rebuild_all_sources(courses: List[str]) -> None:
205
+ """Rebuild all_sources_data.jsonl from active source JSONLs.
206
 
207
  Unlike the previous append-style merge, this drops any source whose entry
208
+ has been removed from source_registry.py and reloads every other active
209
+ source from its own JSONL.
 
210
  """
211
  from data.scraping_scripts.process_md_files import combine_all_sources
212
 
 
410
 
411
 
412
  def update_ui_files(course_name: str) -> None:
413
+ """Confirm the course is represented in the central source registry."""
414
+ if course_name not in SOURCE_KEY_TO_LABEL:
415
+ logger.warning(
416
+ "%s is not in source_registry.py UI metadata. Add it there if it "
417
+ "should appear in the app source picker.",
418
+ course_name,
419
+ )
 
420
  return
421
 
422
+ logger.info(
423
+ "%s is configured in source_registry.py; no setup.py/main.py edits needed.",
424
+ course_name,
425
+ )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
426
 
427
 
428
  def main():
 
433
  "--courses",
434
  nargs="+",
435
  required=True,
436
+ help="One or more course names to process (must match source_registry.py)",
437
  )
438
  parser.add_argument(
439
  "--purge-sources",
 
482
  args = parser.parse_args()
483
  courses: List[str] = args.courses
484
 
 
 
 
 
 
 
 
485
  # Validate every course up front so we fail fast
486
  for course_name in courses:
487
  if course_name not in SOURCE_CONFIGS:
488
  logger.error(f"Course {course_name} not found in SOURCE_CONFIGS")
489
  sys.exit(1)
490
 
491
+ ensure_hf_access()
492
+
493
+ # Keep untouched source JSONLs by downloading them when needed, but don't
494
+ # require first-time courses that this run is about to regenerate.
495
+ sources_to_regenerate = [] if args.skip_process_md else courses
496
+ ensure_required_files_exist(sources_to_regenerate=sources_to_regenerate)
497
+
498
  # Per-course: process markdown + manual URL addition
499
  for course_name in courses:
500
  course_jsonl_path = SOURCE_CONFIGS[course_name]["output_file"]
 
520
  if not args.skip_context:
521
  add_context_to_nodes(not args.process_all_context)
522
 
523
+ prune_contextual_nodes_to_active_sources()
524
+
525
  if not args.skip_vectors:
526
  create_vector_stores()
527
 
data/scraping_scripts/capture_source_versions.py CHANGED
@@ -1,5 +1,5 @@
1
  """
2
- Capture the latest release tag + commit SHA for each documentation source and
3
  write them to `data/source_versions.json`, so the frontend can surface which
4
  library version is represented in the knowledge base (and how fresh it is).
5
 
@@ -7,8 +7,8 @@ Usage:
7
  uv run -m data.scraping_scripts.capture_source_versions
8
  uv run -m data.scraping_scripts.capture_source_versions --sources langchain peft
9
 
10
- This is called automatically from `update_docs_workflow.py` after the GitHub
11
- download step, but can also be run standalone to refresh the JSON.
12
 
13
  Output shape (per source):
14
  {
@@ -31,6 +31,11 @@ from typing import Optional
31
  import requests
32
  from dotenv import load_dotenv
33
 
 
 
 
 
 
34
  load_dotenv()
35
 
36
  OUTPUT_PATH = Path("data/source_versions.json")
@@ -43,6 +48,8 @@ VERSION_REPOS: dict[str, tuple[str, str]] = {
43
  "trl": ("huggingface", "trl"),
44
  "llama_index": ("run-llama", "llama_index"),
45
  "langchain": ("langchain-ai", "langchain"),
 
 
46
  "openai_cookbooks": ("openai", "openai-cookbook"),
47
  }
48
 
@@ -91,7 +98,22 @@ def fetch_default_branch_sha(owner: str, repo: str) -> Optional[str]:
91
  return sha[:7] if isinstance(sha, str) else None
92
 
93
 
94
- def capture_for_source(source: str) -> dict:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
95
  owner, repo = VERSION_REPOS[source]
96
  print(f"- {source}: querying {owner}/{repo}")
97
  return {
@@ -101,6 +123,25 @@ def capture_for_source(source: str) -> dict:
101
  }
102
 
103
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
104
  def load_existing() -> dict:
105
  if not OUTPUT_PATH.exists():
106
  return {}
@@ -115,10 +156,11 @@ def load_existing() -> dict:
115
  def capture(sources: list[str]) -> dict:
116
  versions = load_existing()
117
  for source in sources:
118
- if source not in VERSION_REPOS:
 
119
  print(f" ! skipping unknown source: {source}", file=sys.stderr)
120
  continue
121
- versions[source] = capture_for_source(source)
122
 
123
  OUTPUT_PATH.parent.mkdir(parents=True, exist_ok=True)
124
  with OUTPUT_PATH.open("w", encoding="utf-8") as f:
@@ -133,8 +175,8 @@ def main() -> None:
133
  parser.add_argument(
134
  "--sources",
135
  nargs="+",
136
- choices=list(VERSION_REPOS.keys()),
137
- default=list(VERSION_REPOS.keys()),
138
  help="Subset of sources to refresh (default: all).",
139
  )
140
  args = parser.parse_args()
 
1
  """
2
+ Capture the latest release tag + commit SHA or docs timestamp for each source and
3
  write them to `data/source_versions.json`, so the frontend can surface which
4
  library version is represented in the knowledge base (and how fresh it is).
5
 
 
7
  uv run -m data.scraping_scripts.capture_source_versions
8
  uv run -m data.scraping_scripts.capture_source_versions --sources langchain peft
9
 
10
+ This is called automatically from `update_docs_workflow.py` after the download
11
+ step, but can also be run standalone to refresh the JSON.
12
 
13
  Output shape (per source):
14
  {
 
31
  import requests
32
  from dotenv import load_dotenv
33
 
34
+ try:
35
+ from data.scraping_scripts.source_registry import DOC_SOURCE_KEYS, SOURCE_CONFIGS
36
+ except ModuleNotFoundError:
37
+ from source_registry import DOC_SOURCE_KEYS, SOURCE_CONFIGS
38
+
39
  load_dotenv()
40
 
41
  OUTPUT_PATH = Path("data/source_versions.json")
 
48
  "trl": ("huggingface", "trl"),
49
  "llama_index": ("run-llama", "llama_index"),
50
  "langchain": ("langchain-ai", "langchain"),
51
+ "langgraph": ("langchain-ai", "langgraph"),
52
+ "deep_agents": ("langchain-ai", "deepagents"),
53
  "openai_cookbooks": ("openai", "openai-cookbook"),
54
  }
55
 
 
98
  return sha[:7] if isinstance(sha, str) else None
99
 
100
 
101
+ def fetch_last_modified(url: str) -> Optional[str]:
102
+ try:
103
+ response = requests.head(url, timeout=30, allow_redirects=True)
104
+ except requests.RequestException as exc:
105
+ print(f" ! network error: {exc}", file=sys.stderr)
106
+ return None
107
+ if not response.ok:
108
+ print(
109
+ f" ! HTTP {response.status_code} for {url}: {response.text[:160]}",
110
+ file=sys.stderr,
111
+ )
112
+ return None
113
+ return response.headers.get("Last-Modified")
114
+
115
+
116
+ def capture_github_source(source: str) -> dict:
117
  owner, repo = VERSION_REPOS[source]
118
  print(f"- {source}: querying {owner}/{repo}")
119
  return {
 
123
  }
124
 
125
 
126
+ def capture_llms_txt_source(source: str) -> dict:
127
+ config = SOURCE_CONFIGS[source]
128
+ urls = [str(url) for url in config.get("llms_txt_urls", [])]
129
+ print(f"- {source}: querying llms.txt metadata")
130
+ return {
131
+ "version": fetch_last_modified(urls[0]) if urls else None,
132
+ "sha": None,
133
+ "indexedAt": datetime.now(timezone.utc).date().isoformat(),
134
+ }
135
+
136
+
137
+ def capture_for_source(source: str) -> dict | None:
138
+ if source in VERSION_REPOS:
139
+ return capture_github_source(source)
140
+ if SOURCE_CONFIGS.get(source, {}).get("llms_txt_urls"):
141
+ return capture_llms_txt_source(source)
142
+ return None
143
+
144
+
145
  def load_existing() -> dict:
146
  if not OUTPUT_PATH.exists():
147
  return {}
 
156
  def capture(sources: list[str]) -> dict:
157
  versions = load_existing()
158
  for source in sources:
159
+ metadata = capture_for_source(source)
160
+ if metadata is None:
161
  print(f" ! skipping unknown source: {source}", file=sys.stderr)
162
  continue
163
+ versions[source] = metadata
164
 
165
  OUTPUT_PATH.parent.mkdir(parents=True, exist_ok=True)
166
  with OUTPUT_PATH.open("w", encoding="utf-8") as f:
 
175
  parser.add_argument(
176
  "--sources",
177
  nargs="+",
178
+ choices=list(DOC_SOURCE_KEYS),
179
+ default=list(DOC_SOURCE_KEYS),
180
  help="Subset of sources to refresh (default: all).",
181
  )
182
  args = parser.parse_args()
data/scraping_scripts/contextual_node_pruning.py ADDED
@@ -0,0 +1,67 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Helpers for keeping contextual-node pickles aligned with active sources."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import logging
6
+ import os
7
+ import pickle
8
+ from collections import Counter
9
+ from typing import Any
10
+
11
+ from data.scraping_scripts.source_registry import (
12
+ ACTIVE_SOURCE_KEYS,
13
+ CONTEXTUAL_NODES_PKL,
14
+ )
15
+ from scripts.chroma_rag import get_chunk_record_source
16
+
17
+ logger = logging.getLogger(__name__)
18
+
19
+
20
+ def prune_contextual_nodes_to_active_sources(
21
+ pkl_path: str = CONTEXTUAL_NODES_PKL,
22
+ ) -> Counter[str]:
23
+ """Remove contextual nodes whose source is not active in source_registry.py."""
24
+ if not os.path.exists(pkl_path):
25
+ logger.info("%s does not exist; no contextual nodes to prune", pkl_path)
26
+ return Counter()
27
+
28
+ with open(pkl_path, "rb") as handle:
29
+ nodes = pickle.load(handle)
30
+
31
+ kept_nodes: list[Any] = []
32
+ removed_counts: Counter[str] = Counter()
33
+ unknown_source = 0
34
+
35
+ for node in nodes:
36
+ try:
37
+ source = get_chunk_record_source(node)
38
+ except Exception:
39
+ unknown_source += 1
40
+ kept_nodes.append(node)
41
+ continue
42
+
43
+ if source in ACTIVE_SOURCE_KEYS:
44
+ kept_nodes.append(node)
45
+ else:
46
+ removed_counts[str(source)] += 1
47
+
48
+ if not removed_counts:
49
+ if unknown_source:
50
+ logger.info(
51
+ "No inactive contextual nodes pruned; kept %s nodes with unknown source",
52
+ unknown_source,
53
+ )
54
+ return removed_counts
55
+
56
+ with open(pkl_path, "wb") as handle:
57
+ pickle.dump(kept_nodes, handle)
58
+
59
+ logger.info(
60
+ "Pruned %s inactive contextual nodes from %s: %s",
61
+ sum(removed_counts.values()),
62
+ pkl_path,
63
+ dict(sorted(removed_counts.items())),
64
+ )
65
+ if unknown_source:
66
+ logger.info("Kept %s contextual nodes with unknown source", unknown_source)
67
+ return removed_counts
data/scraping_scripts/create_vector_stores.py CHANGED
@@ -22,13 +22,22 @@ from dotenv import load_dotenv
22
  from tqdm.auto import tqdm
23
 
24
  try:
 
 
 
 
25
  from data.scraping_scripts.add_context_to_nodes import create_docs, process
26
  from scripts.chroma_rag import (
 
27
  DEFAULT_EMBED_MODEL,
 
 
28
  build_document_dict,
29
  embed_texts,
 
30
  load_jsonl_documents,
31
  normalize_chunk_record,
 
32
  save_document_dict,
33
  )
34
  except ModuleNotFoundError:
@@ -38,13 +47,22 @@ except ModuleNotFoundError:
38
  project_root = Path(__file__).resolve().parents[2]
39
  if str(project_root) not in sys.path:
40
  sys.path.insert(0, str(project_root))
 
 
 
 
41
  from data.scraping_scripts.add_context_to_nodes import create_docs, process
42
  from scripts.chroma_rag import (
 
43
  DEFAULT_EMBED_MODEL,
 
 
44
  build_document_dict,
45
  embed_texts,
 
46
  load_jsonl_documents,
47
  normalize_chunk_record,
 
48
  save_document_dict,
49
  )
50
 
@@ -54,55 +72,32 @@ logging.basicConfig(level=logging.INFO)
54
  logger = logging.getLogger(__name__)
55
 
56
 
57
- SOURCE_CONFIGS = {
58
- "transformers": {
59
- "input_file": "data/transformers_data.jsonl",
60
- "db_name": "chroma-db-transformers",
61
- "document_dict_file": "document_dict_transformers.pkl",
62
- },
63
- "peft": {
64
- "input_file": "data/peft_data.jsonl",
65
- "db_name": "chroma-db-peft",
66
- "document_dict_file": "document_dict_peft.pkl",
67
- },
68
- "trl": {
69
- "input_file": "data/trl_data.jsonl",
70
- "db_name": "chroma-db-trl",
71
- "document_dict_file": "document_dict_trl.pkl",
72
- },
73
- "llama_index": {
74
- "input_file": "data/llama_index_data.jsonl",
75
- "db_name": "chroma-db-llama_index",
76
- "document_dict_file": "document_dict_llama_index.pkl",
77
- },
78
- "openai_cookbooks": {
79
- "input_file": "data/openai_cookbooks_data.jsonl",
80
- "db_name": "chroma-db-openai_cookbooks",
81
- "document_dict_file": "document_dict_openai_cookbooks.pkl",
82
- },
83
- "langchain": {
84
- "input_file": "data/langchain_data.jsonl",
85
- "db_name": "chroma-db-langchain",
86
- "document_dict_file": "document_dict_langchain.pkl",
87
- },
88
- "tai_blog": {
89
- "input_file": "data/tai_blog_data.jsonl",
90
- "db_name": "chroma-db-tai_blog",
91
- "document_dict_file": "document_dict_tai_blog.pkl",
92
- },
93
- "all_sources": {
94
- "input_file": "data/all_sources_data.jsonl",
95
- "db_name": "chroma-db-all_sources",
96
- "document_dict_file": "document_dict_all_sources.pkl",
97
- },
98
- }
99
 
100
 
101
  def load_or_create_chunk_records(source: str) -> list[Any]:
102
  config = SOURCE_CONFIGS[source]
103
  if source == "all_sources" and os.path.exists("data/all_sources_contextual_nodes.pkl"):
104
  with open("data/all_sources_contextual_nodes.pkl", "rb") as handle:
105
- return pickle.load(handle)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
106
 
107
  documents = create_docs(config["input_file"])
108
  return asyncio.run(process(documents))
@@ -113,7 +108,51 @@ def iter_batches(items: Sequence[Any], batch_size: int):
113
  yield items[start : start + batch_size]
114
 
115
 
116
- def process_source(source: str) -> None:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
117
  config = SOURCE_CONFIGS[source]
118
  document_rows = load_jsonl_documents(config["input_file"])
119
  if not document_rows:
@@ -122,11 +161,23 @@ def process_source(source: str) -> None:
122
 
123
  db_name = config["db_name"]
124
  db_path = f"data/{db_name}"
125
- if os.path.exists(db_path):
126
  shutil.rmtree(db_path)
127
 
128
  os.makedirs(db_path, exist_ok=True)
129
 
 
 
 
 
 
 
 
 
 
 
 
 
130
  chunk_records = [
131
  normalize_chunk_record(record)
132
  for record in load_or_create_chunk_records(source)
@@ -136,9 +187,6 @@ def process_source(source: str) -> None:
136
  return
137
 
138
  chunk_ids = [record.chunk_id for record in chunk_records]
139
- chunk_texts = [record.text for record in chunk_records]
140
- chunk_metadatas = [record.metadata for record in chunk_records]
141
-
142
  logger.info(
143
  "Preparing %s chunks from %s documents for %s",
144
  len(chunk_records),
@@ -146,58 +194,108 @@ def process_source(source: str) -> None:
146
  source,
147
  )
148
 
149
- cohere_client = cohere.ClientV2(api_key=os.environ["COHERE_API_KEY"])
150
- logger.info("Generating embeddings for %s", source)
151
- embeddings = embed_texts(
152
- cohere_client,
153
- chunk_texts,
154
- input_type="search_document",
155
- model=DEFAULT_EMBED_MODEL,
156
- show_progress=True,
157
- progress_desc=f"Embedding {source}",
158
- )
159
-
160
  chroma_client = chromadb.PersistentClient(path=db_path)
161
  collection = chroma_client.get_or_create_collection(name=db_name)
162
  max_batch_size = chroma_client.get_max_batch_size()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
163
  logger.info(
164
- "Writing %s embeddings to Chroma for %s in batches of up to %s",
165
- len(chunk_ids),
 
166
  source,
167
- max_batch_size,
168
- )
169
- with tqdm(total=len(chunk_ids), desc=f"Upserting {source}", unit="chunk") as progress:
170
- for batch_ids, batch_embeddings, batch_texts, batch_metadatas in zip(
171
- iter_batches(chunk_ids, max_batch_size),
172
- iter_batches(embeddings, max_batch_size),
173
- iter_batches(chunk_texts, max_batch_size),
174
- iter_batches(chunk_metadatas, max_batch_size),
175
- ):
176
- collection.upsert(
177
- ids=batch_ids,
178
- embeddings=batch_embeddings,
179
- documents=batch_texts,
180
- metadatas=batch_metadatas,
181
- )
182
- progress.update(len(batch_ids))
183
 
184
- document_dict = build_document_dict(document_rows)
185
- save_document_dict(
186
- document_dict,
187
- f"{db_path}/{config['document_dict_file']}",
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
188
  )
189
 
190
  print(
191
- f"Indexed {len(chunk_records)} chunks from {len(document_rows)} documents into {db_path}"
 
192
  )
193
 
194
 
195
- def main(sources: list[str]) -> None:
 
 
 
 
 
 
 
 
 
196
  for source in sources:
197
  if source not in SOURCE_CONFIGS:
198
  print(f"Unknown source: {source}")
199
  continue
200
- process_source(source)
 
 
 
 
 
 
 
 
201
 
202
 
203
  if __name__ == "__main__":
@@ -210,5 +308,59 @@ if __name__ == "__main__":
210
  choices=SOURCE_CONFIGS.keys(),
211
  help="Specify one or more sources to process",
212
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
213
  args = parser.parse_args()
214
- main(args.sources)
 
 
 
 
 
 
 
 
 
22
  from tqdm.auto import tqdm
23
 
24
  try:
25
+ from data.scraping_scripts.source_registry import (
26
+ ACTIVE_SOURCE_KEYS,
27
+ vector_store_source_configs,
28
+ )
29
  from data.scraping_scripts.add_context_to_nodes import create_docs, process
30
  from scripts.chroma_rag import (
31
+ DEFAULT_COHERE_EMBED_BATCH_SIZE,
32
  DEFAULT_EMBED_MODEL,
33
+ BM25Index,
34
+ build_chunk_records,
35
  build_document_dict,
36
  embed_texts,
37
+ get_chunk_record_source,
38
  load_jsonl_documents,
39
  normalize_chunk_record,
40
+ save_bm25_index,
41
  save_document_dict,
42
  )
43
  except ModuleNotFoundError:
 
47
  project_root = Path(__file__).resolve().parents[2]
48
  if str(project_root) not in sys.path:
49
  sys.path.insert(0, str(project_root))
50
+ from data.scraping_scripts.source_registry import (
51
+ ACTIVE_SOURCE_KEYS,
52
+ vector_store_source_configs,
53
+ )
54
  from data.scraping_scripts.add_context_to_nodes import create_docs, process
55
  from scripts.chroma_rag import (
56
+ DEFAULT_COHERE_EMBED_BATCH_SIZE,
57
  DEFAULT_EMBED_MODEL,
58
+ BM25Index,
59
+ build_chunk_records,
60
  build_document_dict,
61
  embed_texts,
62
+ get_chunk_record_source,
63
  load_jsonl_documents,
64
  normalize_chunk_record,
65
+ save_bm25_index,
66
  save_document_dict,
67
  )
68
 
 
72
  logger = logging.getLogger(__name__)
73
 
74
 
75
+ SOURCE_CONFIGS = vector_store_source_configs()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
76
 
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 = []
84
+ skipped_count = 0
85
+ for record in records:
86
+ try:
87
+ record_source = get_chunk_record_source(record)
88
+ except Exception:
89
+ active_records.append(record)
90
+ continue
91
+ if record_source in ACTIVE_SOURCE_KEYS:
92
+ active_records.append(record)
93
+ else:
94
+ skipped_count += 1
95
+ if skipped_count:
96
+ logger.info(
97
+ "Skipped %s inactive contextual chunks while building all_sources",
98
+ skipped_count,
99
+ )
100
+ return active_records
101
 
102
  documents = create_docs(config["input_file"])
103
  return asyncio.run(process(documents))
 
108
  yield items[start : start + batch_size]
109
 
110
 
111
+ def get_collection_ids(collection, batch_size: int = 5000) -> set[str]:
112
+ ids: set[str] = set()
113
+ offset = 0
114
+
115
+ while True:
116
+ result = collection.get(limit=batch_size, offset=offset, include=[])
117
+ batch_ids = result.get("ids", [])
118
+ if not batch_ids:
119
+ break
120
+
121
+ ids.update(str(chunk_id) for chunk_id in batch_ids)
122
+ offset += len(batch_ids)
123
+
124
+ return ids
125
+
126
+
127
+ def write_retrieval_artifacts(
128
+ *,
129
+ config: dict[str, str],
130
+ document_rows: list[dict[str, Any]],
131
+ db_path: str,
132
+ ) -> int:
133
+ document_dict = build_document_dict(document_rows)
134
+ save_document_dict(
135
+ document_dict,
136
+ f"{db_path}/{config['document_dict_file']}",
137
+ )
138
+ bm25_records = build_chunk_records(document_rows)
139
+ save_bm25_index(
140
+ BM25Index.build(bm25_records),
141
+ f"{db_path}/{config['bm25_index_file']}",
142
+ )
143
+ return len(bm25_records)
144
+
145
+
146
+ def process_source(
147
+ source: str,
148
+ *,
149
+ force_rebuild: bool = False,
150
+ skip_dense_embeddings: bool = False,
151
+ embed_batch_size: int = DEFAULT_COHERE_EMBED_BATCH_SIZE,
152
+ cohere_embed_inputs_per_minute: int | None = None,
153
+ cohere_embed_tpm_limit: int | None = None,
154
+ cohere_embed_rpm_limit: int | None = None,
155
+ ) -> None:
156
  config = SOURCE_CONFIGS[source]
157
  document_rows = load_jsonl_documents(config["input_file"])
158
  if not document_rows:
 
161
 
162
  db_name = config["db_name"]
163
  db_path = f"data/{db_name}"
164
+ if force_rebuild and not skip_dense_embeddings and os.path.exists(db_path):
165
  shutil.rmtree(db_path)
166
 
167
  os.makedirs(db_path, exist_ok=True)
168
 
169
+ if skip_dense_embeddings:
170
+ bm25_count = write_retrieval_artifacts(
171
+ config=config,
172
+ document_rows=document_rows,
173
+ db_path=db_path,
174
+ )
175
+ print(
176
+ f"Indexed {bm25_count} BM25 chunks from {len(document_rows)} documents "
177
+ f"into {db_path}; skipped dense embedding updates"
178
+ )
179
+ return
180
+
181
  chunk_records = [
182
  normalize_chunk_record(record)
183
  for record in load_or_create_chunk_records(source)
 
187
  return
188
 
189
  chunk_ids = [record.chunk_id for record in chunk_records]
 
 
 
190
  logger.info(
191
  "Preparing %s chunks from %s documents for %s",
192
  len(chunk_records),
 
194
  source,
195
  )
196
 
 
 
 
 
 
 
 
 
 
 
 
197
  chroma_client = chromadb.PersistentClient(path=db_path)
198
  collection = chroma_client.get_or_create_collection(name=db_name)
199
  max_batch_size = chroma_client.get_max_batch_size()
200
+ existing_ids = set() if force_rebuild else get_collection_ids(collection)
201
+ desired_ids = set(chunk_ids)
202
+
203
+ stale_ids = sorted(existing_ids - desired_ids)
204
+ if stale_ids:
205
+ logger.info("Deleting %s stale embeddings for %s", len(stale_ids), source)
206
+ for batch_ids in iter_batches(stale_ids, max_batch_size):
207
+ collection.delete(ids=batch_ids)
208
+
209
+ reusable_ids = existing_ids & desired_ids
210
+ records_to_embed = [
211
+ record for record in chunk_records if record.chunk_id not in reusable_ids
212
+ ]
213
+
214
  logger.info(
215
+ "Reusing %s existing embeddings and generating %s embeddings for %s",
216
+ len(reusable_ids),
217
+ len(records_to_embed),
218
  source,
219
+ )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
220
 
221
+ if records_to_embed:
222
+ cohere_client = cohere.ClientV2(api_key=os.environ["COHERE_API_KEY"])
223
+ texts_to_embed = [record.text for record in records_to_embed]
224
+ ids_to_embed = [record.chunk_id for record in records_to_embed]
225
+ metadatas_to_embed = [record.metadata for record in records_to_embed]
226
+
227
+ logger.info("Generating embeddings for %s", source)
228
+ embeddings = embed_texts(
229
+ cohere_client,
230
+ texts_to_embed,
231
+ input_type="search_document",
232
+ model=DEFAULT_EMBED_MODEL,
233
+ batch_size=embed_batch_size,
234
+ max_inputs_per_minute=cohere_embed_inputs_per_minute,
235
+ max_tokens_per_minute=cohere_embed_tpm_limit,
236
+ max_requests_per_minute=cohere_embed_rpm_limit,
237
+ show_progress=True,
238
+ progress_desc=f"Embedding {source}",
239
+ )
240
+
241
+ logger.info(
242
+ "Writing %s new embeddings to Chroma for %s in batches of up to %s",
243
+ len(ids_to_embed),
244
+ source,
245
+ max_batch_size,
246
+ )
247
+ with tqdm(
248
+ total=len(ids_to_embed), desc=f"Upserting {source}", unit="chunk"
249
+ ) as progress:
250
+ for batch_ids, batch_embeddings, batch_texts, batch_metadatas in zip(
251
+ iter_batches(ids_to_embed, max_batch_size),
252
+ iter_batches(embeddings, max_batch_size),
253
+ iter_batches(texts_to_embed, max_batch_size),
254
+ iter_batches(metadatas_to_embed, max_batch_size),
255
+ ):
256
+ collection.upsert(
257
+ ids=batch_ids,
258
+ embeddings=batch_embeddings,
259
+ documents=batch_texts,
260
+ metadatas=batch_metadatas,
261
+ )
262
+ progress.update(len(batch_ids))
263
+
264
+ bm25_count = write_retrieval_artifacts(
265
+ config=config,
266
+ document_rows=document_rows,
267
+ db_path=db_path,
268
  )
269
 
270
  print(
271
+ f"Indexed {len(chunk_records)} dense chunks and {bm25_count} BM25 chunks "
272
+ f"from {len(document_rows)} documents into {db_path}"
273
  )
274
 
275
 
276
+ def main(
277
+ sources: list[str],
278
+ *,
279
+ force_rebuild: bool = False,
280
+ skip_dense_embeddings: bool = False,
281
+ embed_batch_size: int = DEFAULT_COHERE_EMBED_BATCH_SIZE,
282
+ cohere_embed_inputs_per_minute: int | None = None,
283
+ cohere_embed_tpm_limit: int | None = None,
284
+ cohere_embed_rpm_limit: int | None = None,
285
+ ) -> None:
286
  for source in sources:
287
  if source not in SOURCE_CONFIGS:
288
  print(f"Unknown source: {source}")
289
  continue
290
+ process_source(
291
+ source,
292
+ force_rebuild=force_rebuild,
293
+ skip_dense_embeddings=skip_dense_embeddings,
294
+ embed_batch_size=embed_batch_size,
295
+ cohere_embed_inputs_per_minute=cohere_embed_inputs_per_minute,
296
+ cohere_embed_tpm_limit=cohere_embed_tpm_limit,
297
+ cohere_embed_rpm_limit=cohere_embed_rpm_limit,
298
+ )
299
 
300
 
301
  if __name__ == "__main__":
 
308
  choices=SOURCE_CONFIGS.keys(),
309
  help="Specify one or more sources to process",
310
  )
311
+ parser.add_argument(
312
+ "--force-rebuild",
313
+ action="store_true",
314
+ help="Delete the existing Chroma directory and regenerate all embeddings.",
315
+ )
316
+ parser.add_argument(
317
+ "--skip-dense-embeddings",
318
+ action="store_true",
319
+ help=(
320
+ "Only write retrieval artifacts that do not require Cohere embeddings "
321
+ "(document dictionary and BM25 index). Existing Chroma data is left untouched."
322
+ ),
323
+ )
324
+ parser.add_argument(
325
+ "--embed-batch-size",
326
+ type=int,
327
+ default=DEFAULT_COHERE_EMBED_BATCH_SIZE,
328
+ help="Maximum number of chunks to include in one Cohere embed request.",
329
+ )
330
+ parser.add_argument(
331
+ "--cohere-embed-inputs-per-minute",
332
+ type=int,
333
+ default=None,
334
+ help=(
335
+ "Cohere embed input-per-minute limit before the safety margin. "
336
+ "Defaults to COHERE_EMBED_INPUTS_PER_MINUTE or 2000; use 0 to disable."
337
+ ),
338
+ )
339
+ parser.add_argument(
340
+ "--cohere-embed-tpm-limit",
341
+ type=int,
342
+ default=None,
343
+ help=(
344
+ "Cohere embed token-per-minute limit before the safety margin. "
345
+ "Defaults to COHERE_EMBED_TPM_LIMIT or 0; use 0 to disable."
346
+ ),
347
+ )
348
+ parser.add_argument(
349
+ "--cohere-embed-rpm-limit",
350
+ type=int,
351
+ default=None,
352
+ help=(
353
+ "Cohere embed request-per-minute limit before the safety margin. "
354
+ "Defaults to COHERE_EMBED_RPM_LIMIT or 0; use 0 to disable."
355
+ ),
356
+ )
357
  args = parser.parse_args()
358
+ main(
359
+ args.sources,
360
+ force_rebuild=args.force_rebuild,
361
+ skip_dense_embeddings=args.skip_dense_embeddings,
362
+ embed_batch_size=args.embed_batch_size,
363
+ cohere_embed_inputs_per_minute=args.cohere_embed_inputs_per_minute,
364
+ cohere_embed_tpm_limit=args.cohere_embed_tpm_limit,
365
+ cohere_embed_rpm_limit=args.cohere_embed_rpm_limit,
366
+ )
data/scraping_scripts/github_to_markdown_ai_docs.py CHANGED
@@ -42,6 +42,11 @@ import requests
42
  from dotenv import load_dotenv
43
  from nbconvert import MarkdownExporter
44
 
 
 
 
 
 
45
  load_dotenv()
46
 
47
  # Configuration for different sources
@@ -113,6 +118,23 @@ SOURCE_CONFIGS = {
113
  ],
114
  "local_dir": "data/langchain_md_files",
115
  },
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
116
  }
117
 
118
  # GitHub Personal Access Token (replace with your own token)
@@ -125,6 +147,7 @@ if GITHUB_TOKEN:
125
 
126
  # Maximum number of retries
127
  MAX_RETRIES = 5
 
128
 
129
 
130
  class GitHubAPIError(RuntimeError):
@@ -242,11 +265,22 @@ def download_file(file_url: str, file_path: str, retries: int = 0):
242
  # f.write(markdown)
243
 
244
 
 
 
 
 
 
 
 
 
 
 
245
  def convert_ipynb_to_md(ipynb_path: str, md_path: str):
246
  try:
247
  with open(ipynb_path, "r", encoding="utf-8") as f:
248
  notebook = nbformat.read(f, as_version=4)
249
 
 
250
  exporter = MarkdownExporter()
251
  markdown, _ = exporter.from_notebook_node(notebook)
252
 
@@ -260,7 +294,23 @@ def convert_ipynb_to_md(ipynb_path: str, md_path: str):
260
  print("Skipping this file and continuing with others...")
261
 
262
 
263
- def fetch_files(api_url: str, local_dir: str):
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
264
  files = get_files_in_directory(api_url)
265
  if isinstance(files, dict):
266
  files = [files]
@@ -278,10 +328,14 @@ def fetch_files(api_url: str, local_dir: str):
278
  print(f"Converting {file_name} to markdown...")
279
  convert_ipynb_to_md(file_path, md_file_path)
280
  os.remove(file_path) # Remove the .ipynb file after conversion
 
 
 
 
281
  elif file["type"] == "dir":
282
  subdir = os.path.join(local_dir, file["name"])
283
  os.makedirs(subdir, exist_ok=True)
284
- fetch_files(file["url"], subdir)
285
 
286
 
287
  def process_source(source: str):
@@ -295,6 +349,7 @@ def process_source(source: str):
295
  local_dir = config.get("local_dir", f"data/{config['repo']}_md_files")
296
  shutil.rmtree(local_dir, ignore_errors=True)
297
  os.makedirs(local_dir, exist_ok=True)
 
298
 
299
  print(f"Processing source: {source}")
300
 
@@ -306,13 +361,14 @@ def process_source(source: str):
306
  )
307
  target_dir = os.path.join(local_dir, path_config.get("local_subdir", ""))
308
  os.makedirs(target_dir, exist_ok=True)
309
- fetch_files(api_url, target_dir)
310
  else:
311
  api_url = (
312
  f"https://api.github.com/repos/{config['owner']}/{config['repo']}/contents/{config['path']}"
313
  )
314
- fetch_files(api_url, local_dir)
315
 
 
316
  print(f"Finished processing {source}")
317
 
318
 
 
42
  from dotenv import load_dotenv
43
  from nbconvert import MarkdownExporter
44
 
45
+ try:
46
+ from data.scraping_scripts.source_registry import GITHUB_SOURCE_KEYS
47
+ except ModuleNotFoundError:
48
+ from source_registry import GITHUB_SOURCE_KEYS
49
+
50
  load_dotenv()
51
 
52
  # Configuration for different sources
 
118
  ],
119
  "local_dir": "data/langchain_md_files",
120
  },
121
+ "langgraph": {
122
+ "owner": "langchain-ai",
123
+ "repo": "docs",
124
+ "path": "src/oss/langgraph",
125
+ "local_dir": "data/langgraph_md_files",
126
+ },
127
+ "deep_agents": {
128
+ "owner": "langchain-ai",
129
+ "repo": "docs",
130
+ "path": "src/oss/deepagents",
131
+ "local_dir": "data/deep_agents_md_files",
132
+ },
133
+ }
134
+ SOURCE_CONFIGS = {
135
+ source: config
136
+ for source, config in SOURCE_CONFIGS.items()
137
+ if source in GITHUB_SOURCE_KEYS
138
  }
139
 
140
  # GitHub Personal Access Token (replace with your own token)
 
147
 
148
  # Maximum number of retries
149
  MAX_RETRIES = 5
150
+ SOURCE_EXTENSIONS_FILENAME = "_source_extensions.json"
151
 
152
 
153
  class GitHubAPIError(RuntimeError):
 
265
  # f.write(markdown)
266
 
267
 
268
+ def strip_notebook_outputs(notebook: nbformat.NotebookNode) -> nbformat.NotebookNode:
269
+ for cell in notebook.cells:
270
+ if cell.get("cell_type") == "code":
271
+ cell["outputs"] = []
272
+ cell["execution_count"] = None
273
+ if "attachments" in cell:
274
+ cell["attachments"] = {}
275
+ return notebook
276
+
277
+
278
  def convert_ipynb_to_md(ipynb_path: str, md_path: str):
279
  try:
280
  with open(ipynb_path, "r", encoding="utf-8") as f:
281
  notebook = nbformat.read(f, as_version=4)
282
 
283
+ notebook = strip_notebook_outputs(notebook)
284
  exporter = MarkdownExporter()
285
  markdown, _ = exporter.from_notebook_node(notebook)
286
 
 
294
  print("Skipping this file and continuing with others...")
295
 
296
 
297
+ def manifest_path(file_path: str, root_dir: str) -> str:
298
+ return os.path.relpath(file_path, root_dir).replace(os.sep, "/")
299
+
300
+
301
+ def save_source_extension_manifest(local_dir: str, manifest: dict[str, str]) -> None:
302
+ path = os.path.join(local_dir, SOURCE_EXTENSIONS_FILENAME)
303
+ with open(path, "w", encoding="utf-8") as f:
304
+ json.dump(dict(sorted(manifest.items())), f, indent=2)
305
+ f.write("\n")
306
+
307
+
308
+ def fetch_files(
309
+ api_url: str,
310
+ local_dir: str,
311
+ source_extensions: dict[str, str],
312
+ root_dir: str,
313
+ ):
314
  files = get_files_in_directory(api_url)
315
  if isinstance(files, dict):
316
  files = [files]
 
328
  print(f"Converting {file_name} to markdown...")
329
  convert_ipynb_to_md(file_path, md_file_path)
330
  os.remove(file_path) # Remove the .ipynb file after conversion
331
+ source_extensions[manifest_path(md_file_path, root_dir)] = ".ipynb"
332
+ else:
333
+ _, extension = os.path.splitext(file_name)
334
+ source_extensions[manifest_path(file_path, root_dir)] = extension
335
  elif file["type"] == "dir":
336
  subdir = os.path.join(local_dir, file["name"])
337
  os.makedirs(subdir, exist_ok=True)
338
+ fetch_files(file["url"], subdir, source_extensions, root_dir)
339
 
340
 
341
  def process_source(source: str):
 
349
  local_dir = config.get("local_dir", f"data/{config['repo']}_md_files")
350
  shutil.rmtree(local_dir, ignore_errors=True)
351
  os.makedirs(local_dir, exist_ok=True)
352
+ source_extensions: dict[str, str] = {}
353
 
354
  print(f"Processing source: {source}")
355
 
 
361
  )
362
  target_dir = os.path.join(local_dir, path_config.get("local_subdir", ""))
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)
372
  print(f"Finished processing {source}")
373
 
374
 
data/scraping_scripts/llms_txt_to_markdown_docs.py ADDED
@@ -0,0 +1,194 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Fetch Markdown documentation pages listed in llms.txt indexes.
3
+
4
+ This downloader is for official docs sites that expose AI-friendly Markdown
5
+ indexes, such as OpenAI's developer docs. It writes the linked Markdown pages
6
+ into a local directory so the existing process_md_files.py pipeline can create
7
+ JSONL, contextual nodes, and vector stores without needing an HTML crawler.
8
+
9
+ Usage:
10
+ uv run -m data.scraping_scripts.llms_txt_to_markdown_docs openai_docs
11
+ uv run -m data.scraping_scripts.llms_txt_to_markdown_docs openai_docs --dry-run
12
+ """
13
+
14
+ from __future__ import annotations
15
+
16
+ import argparse
17
+ import json
18
+ import os
19
+ import re
20
+ import shutil
21
+ from datetime import datetime, timezone
22
+ from typing import Iterable
23
+ from urllib.parse import urlparse
24
+
25
+ import requests
26
+ from dotenv import load_dotenv
27
+
28
+ try:
29
+ from data.scraping_scripts.source_registry import (
30
+ LLMS_TXT_SOURCE_KEYS,
31
+ SOURCE_CONFIGS,
32
+ )
33
+ except ModuleNotFoundError:
34
+ from source_registry import LLMS_TXT_SOURCE_KEYS, SOURCE_CONFIGS
35
+
36
+ load_dotenv()
37
+
38
+ LINK_PATTERN = re.compile(r"\[[^\]]+\]\((https?://[^)\s]+)\)")
39
+ SOURCE_EXTENSIONS_FILENAME = "_source_extensions.json"
40
+ SOURCE_URLS_FILENAME = "_source_urls.json"
41
+ LLMS_TXT_MANIFEST_FILENAME = "_llms_txt_manifest.json"
42
+ REQUEST_TIMEOUT_SECONDS = 60
43
+
44
+
45
+ class LLMSTxtDownloadError(RuntimeError):
46
+ """Raised when an llms.txt source cannot be downloaded safely."""
47
+
48
+
49
+ def fetch_text(url: str) -> str:
50
+ response = requests.get(url, timeout=REQUEST_TIMEOUT_SECONDS)
51
+ response.raise_for_status()
52
+ return response.text
53
+
54
+
55
+ def parse_markdown_links(index_text: str) -> list[str]:
56
+ """Extract HTTP(S) Markdown links from an llms.txt index."""
57
+ urls = []
58
+ seen = set()
59
+ for match in LINK_PATTERN.finditer(index_text):
60
+ url = match.group(1).strip()
61
+ if url not in seen:
62
+ urls.append(url)
63
+ seen.add(url)
64
+ return urls
65
+
66
+
67
+ def should_download_url(url: str, include_prefixes: Iterable[str]) -> bool:
68
+ parsed = urlparse(url)
69
+ if parsed.scheme not in {"http", "https"}:
70
+ return False
71
+ if not parsed.path.endswith((".md", ".mdx")):
72
+ return False
73
+
74
+ prefixes = tuple(include_prefixes)
75
+ return not prefixes or any(url.startswith(prefix) for prefix in prefixes)
76
+
77
+
78
+ def local_relative_path(url: str) -> str:
79
+ parsed = urlparse(url)
80
+ relative_path = parsed.path.lstrip("/")
81
+ if not relative_path:
82
+ raise LLMSTxtDownloadError(f"Could not derive a local path from URL: {url}")
83
+ normalized = os.path.normpath(relative_path).replace(os.sep, "/")
84
+ if normalized == ".." or normalized.startswith("../"):
85
+ raise LLMSTxtDownloadError(f"Refusing unsafe local path for URL: {url}")
86
+ return normalized
87
+
88
+
89
+ def write_text_file(path: str, content: str) -> None:
90
+ os.makedirs(os.path.dirname(path), exist_ok=True)
91
+ with open(path, "w", encoding="utf-8") as f:
92
+ f.write(content)
93
+ if content and not content.endswith("\n"):
94
+ f.write("\n")
95
+
96
+
97
+ def write_json(path: str, payload: object) -> None:
98
+ with open(path, "w", encoding="utf-8") as f:
99
+ json.dump(payload, f, indent=2, sort_keys=True)
100
+ f.write("\n")
101
+
102
+
103
+ def collect_doc_urls(config: dict) -> list[str]:
104
+ index_urls = config.get("llms_txt_urls") or []
105
+ if not index_urls:
106
+ raise LLMSTxtDownloadError("Source config is missing llms_txt_urls")
107
+
108
+ include_prefixes = config.get("llms_url_include_prefixes") or []
109
+ docs: list[str] = []
110
+ seen = set()
111
+
112
+ for index_url in index_urls:
113
+ print(f"Fetching index: {index_url}")
114
+ index_text = fetch_text(str(index_url))
115
+ for url in parse_markdown_links(index_text):
116
+ if not should_download_url(url, include_prefixes):
117
+ continue
118
+ if url in seen:
119
+ continue
120
+ docs.append(url)
121
+ seen.add(url)
122
+
123
+ return docs
124
+
125
+
126
+ def process_source(source: str, *, dry_run: bool = False) -> None:
127
+ if source not in LLMS_TXT_SOURCE_KEYS:
128
+ available = ", ".join(LLMS_TXT_SOURCE_KEYS)
129
+ raise LLMSTxtDownloadError(
130
+ f"Unknown llms.txt source '{source}'. Available sources: {available}"
131
+ )
132
+
133
+ config = SOURCE_CONFIGS[source]
134
+ doc_urls = collect_doc_urls(config)
135
+
136
+ if not doc_urls:
137
+ raise LLMSTxtDownloadError(f"No Markdown docs found for source: {source}")
138
+
139
+ if dry_run:
140
+ print(f"Found {len(doc_urls)} Markdown docs for {source}")
141
+ return
142
+
143
+ local_dir = str(config["input_directory"])
144
+ shutil.rmtree(local_dir, ignore_errors=True)
145
+ os.makedirs(local_dir, exist_ok=True)
146
+
147
+ source_extensions: dict[str, str] = {}
148
+ source_urls: dict[str, str] = {}
149
+
150
+ print(f"Downloading {len(doc_urls)} Markdown docs for {source}")
151
+ for doc_url in doc_urls:
152
+ relative_path = local_relative_path(doc_url)
153
+ local_path = os.path.join(local_dir, relative_path)
154
+ print(f"Downloading {doc_url}")
155
+ content = fetch_text(doc_url)
156
+ write_text_file(local_path, content)
157
+ source_extensions[relative_path] = os.path.splitext(relative_path)[1]
158
+ source_urls[relative_path] = doc_url
159
+
160
+ write_json(os.path.join(local_dir, SOURCE_EXTENSIONS_FILENAME), source_extensions)
161
+ write_json(os.path.join(local_dir, SOURCE_URLS_FILENAME), source_urls)
162
+ write_json(
163
+ os.path.join(local_dir, LLMS_TXT_MANIFEST_FILENAME),
164
+ {
165
+ "downloadedAt": datetime.now(timezone.utc).isoformat(),
166
+ "indexUrls": config.get("llms_txt_urls") or [],
167
+ "documentCount": len(doc_urls),
168
+ },
169
+ )
170
+ print(f"Finished processing {source}")
171
+
172
+
173
+ def main(sources: list[str], *, dry_run: bool = False) -> None:
174
+ for source in sources:
175
+ process_source(source, dry_run=dry_run)
176
+ print("All specified llms.txt sources have been processed.")
177
+
178
+
179
+ if __name__ == "__main__":
180
+ parser = argparse.ArgumentParser(description=__doc__)
181
+ parser.add_argument(
182
+ "sources",
183
+ nargs="+",
184
+ choices=LLMS_TXT_SOURCE_KEYS,
185
+ help="Specify one or more llms.txt-backed sources to process",
186
+ )
187
+ parser.add_argument(
188
+ "--dry-run",
189
+ action="store_true",
190
+ help="Fetch indexes and report document counts without writing files",
191
+ )
192
+ args = parser.parse_args()
193
+
194
+ main(args.sources, dry_run=args.dry_run)
data/scraping_scripts/process_md_files.py CHANGED
@@ -18,14 +18,15 @@ Key features:
18
  Usage:
19
  python process_md_files.py <source1> <source2> ...
20
 
21
- Where <source1>, <source2>, etc. are one or more of the predefined sources in SOURCE_CONFIGS
 
22
  (e.g., 'transformers', 'llama_index', 'openai_cookbooks').
23
 
24
  The script processes all Markdown files in the specified input directories (and their subdirectories),
25
  applies the configured filters, and saves the results in JSONL files. Each line in the output
26
  files represents a single document with metadata and content.
27
 
28
- To add or modify sources, update the SOURCE_CONFIGS dictionary at the top of the script.
29
  """
30
 
31
  import argparse
@@ -38,170 +39,90 @@ from typing import Dict, List
38
 
39
  import tiktoken
40
 
 
 
 
 
 
41
  logging.basicConfig(level=logging.INFO)
42
  logger = logging.getLogger(__name__)
43
 
44
- # Configuration for different sources
45
- SOURCE_CONFIGS = {
46
- "transformers": {
47
- "base_url": "https://huggingface.co/docs/transformers/",
48
- "input_directory": "data/transformers_md_files",
49
- "output_file": "data/transformers_data.jsonl",
50
- "source_name": "transformers",
51
- "use_include_list": False,
52
- "included_dirs": [],
53
- "excluded_dirs": ["internal", "main_classes"],
54
- "excluded_root_files": [],
55
- "included_root_files": [],
56
- "url_extension": "",
57
- },
58
- "peft": {
59
- "base_url": "https://huggingface.co/docs/peft/",
60
- "input_directory": "data/peft_md_files",
61
- "output_file": "data/peft_data.jsonl",
62
- "source_name": "peft",
63
- "use_include_list": False,
64
- "included_dirs": [],
65
- "excluded_dirs": [],
66
- "excluded_root_files": [],
67
- "included_root_files": [],
68
- "url_extension": "",
69
- },
70
- "trl": {
71
- "base_url": "https://huggingface.co/docs/trl/",
72
- "input_directory": "data/trl_md_files",
73
- "output_file": "data/trl_data.jsonl",
74
- "source_name": "trl",
75
- "use_include_list": False,
76
- "included_dirs": [],
77
- "excluded_dirs": [],
78
- "excluded_root_files": [],
79
- "included_root_files": [],
80
- "url_extension": "",
81
- },
82
- "llama_index": {
83
- "base_url": "https://docs.llamaindex.ai/en/stable/",
84
- "input_directory": "data/llama_index_md_files",
85
- "output_file": "data/llama_index_data.jsonl",
86
- "source_name": "llama_index",
87
- "use_include_list": True,
88
- "included_dirs": [
89
- "src/content/docs/framework/index.md",
90
- "src/content/docs/framework/getting_started",
91
- "src/content/docs/framework/understanding",
92
- "src/content/docs/framework/use_cases",
93
- "src/content/docs/framework/module_guides",
94
- "src/content/docs/framework/optimizing",
95
- "examples",
96
- ],
97
- "excluded_dirs": [],
98
- "excluded_root_files": [],
99
- "included_root_files": [],
100
- "url_extension": "",
101
- },
102
- "openai_cookbooks": {
103
- "base_url": "https://github.com/openai/openai-cookbook/blob/main/examples/",
104
- "input_directory": "data/openai-cookbook_md_files",
105
- "output_file": "data/openai_cookbooks_data.jsonl",
106
- "source_name": "openai_cookbooks",
107
- "use_include_list": False,
108
- "included_dirs": [],
109
- "excluded_dirs": [],
110
- "excluded_root_files": [],
111
- "included_root_files": [],
112
- "url_extension": ".ipynb",
113
- },
114
- "langchain": {
115
- "base_url": "https://docs.langchain.com/oss/python/",
116
- "input_directory": "data/langchain_md_files",
117
- "output_file": "data/langchain_data.jsonl",
118
- "source_name": "langchain",
119
- "use_include_list": True,
120
- "included_dirs": [
121
- "concepts",
122
- "langchain",
123
- "python/integrations",
124
- "python/migrate",
125
- "python/releases",
126
- ],
127
- "excluded_dirs": [],
128
- "excluded_root_files": [],
129
- "included_root_files": [
130
- "security-policy.mdx",
131
- "release-policy.mdx",
132
- "versioning.mdx",
133
- ],
134
- "url_extension": "",
135
- },
136
- "tai_blog": {
137
- "base_url": "",
138
- "input_directory": "",
139
- "output_file": "data/tai_blog_data.jsonl",
140
- "source_name": "tai_blog",
141
- "use_include_list": False,
142
- "included_dirs": [],
143
- "excluded_dirs": [],
144
- "excluded_root_files": [],
145
- "included_root_files": [],
146
- "url_extension": "",
147
- },
148
- "full_stack_ai_engineering": {
149
- "base_url": "",
150
- "input_directory": "data/full_stack_ai_engineering", # Path to the directory that contains the Markdown files
151
- "output_file": "data/full_stack_ai_engineering_data.jsonl",
152
- "source_name": "full_stack_ai_engineering",
153
- "use_include_list": False,
154
- "included_dirs": [],
155
- "excluded_dirs": [],
156
- "excluded_root_files": [],
157
- "included_root_files": [],
158
- "url_extension": "",
159
- },
160
- "beginner_python_for_ai_engineering": {
161
- "base_url": "",
162
- "input_directory": "data/beginner_python_for_ai_engineering", # Path to the directory that contains the Markdown files
163
- "output_file": "data/beginner_python_for_ai_engineering_data.jsonl",
164
- "source_name": "beginner_python_for_ai_engineering",
165
- "use_include_list": False,
166
- "included_dirs": [],
167
- "excluded_dirs": [],
168
- "excluded_root_files": [],
169
- "included_root_files": [],
170
- "url_extension": "",
171
- },
172
- "master_ai_for_work": {
173
- "base_url": "",
174
- "input_directory": "data/master_ai_for_work", # Path to the directory that contains the Markdown files
175
- "output_file": "data/master_ai_for_work_data.jsonl",
176
- "source_name": "master_ai_for_work",
177
- "use_include_list": False,
178
- "included_dirs": [],
179
- "excluded_dirs": [],
180
- "excluded_root_files": [],
181
- "included_root_files": [],
182
- "url_extension": "",
183
- },
184
- "agentic_ai_engineering": {
185
- "base_url": "",
186
- "input_directory": "data/agentic_ai_engineering", # Path to the directory that contains the Markdown files
187
- "output_file": "data/agentic_ai_engineering_data.jsonl", # Agentic AI Engineering
188
- "source_name": "agentic_ai_engineering",
189
- "use_include_list": False,
190
- "included_dirs": [],
191
- "excluded_dirs": [],
192
- "excluded_root_files": [],
193
- "included_root_files": [],
194
- "url_extension": "",
195
- },
196
- }
197
 
 
 
 
 
198
 
199
- def extract_title(content: str):
200
- title_match = re.search(r"^#\s+(.+)$", content, re.MULTILINE)
201
- if title_match:
202
- return title_match.group(1).strip()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
203
 
204
- lines = content.split("\n")
 
 
 
 
 
 
 
 
 
 
 
 
 
205
  for line in lines:
206
  if line.strip():
207
  return line.strip()
@@ -209,16 +130,70 @@ def extract_title(content: str):
209
  return None
210
 
211
 
212
- def generate_url(file_path: str, config: Dict) -> str:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
213
  """
214
  Return an empty string if base_url is empty;
215
  otherwise return the constructed URL as before.
216
  """
 
 
 
217
  if not config["base_url"]:
218
  return ""
219
 
220
- path_without_extension = os.path.splitext(file_path)[0]
221
  source_name = config["source_name"]
 
 
 
 
 
 
 
 
 
222
 
223
  if source_name == "llama_index":
224
  framework_prefix = "src/content/docs/framework/"
@@ -261,8 +236,36 @@ def remove_copyright_header(content: str) -> str:
261
  return cleaned_content.strip()
262
 
263
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
264
  def process_md_files(directory: str, config: Dict) -> List[Dict]:
265
  jsonl_data = []
 
 
266
 
267
  for root, _, files in os.walk(directory):
268
  for file in files:
@@ -274,8 +277,9 @@ def process_md_files(directory: str, config: Dict) -> List[Dict]:
274
  with open(file_path, "r", encoding="utf-8") as f:
275
  content = f.read()
276
 
277
- title = extract_title(content)
278
- token_count = num_tokens_from_string(content, "cl100k_base")
 
279
 
280
  # Skip very small or extremely large files
281
  if token_count < 100 or token_count > 200_000:
@@ -284,13 +288,18 @@ def process_md_files(directory: str, config: Dict) -> List[Dict]:
284
  )
285
  continue
286
 
287
- cleaned_content = remove_copyright_header(content)
288
-
289
  json_object = {
290
  "tokens": token_count,
291
  "doc_id": str(uuid.uuid4()),
292
  "name": (title if title else file),
293
- "url": generate_url(relative_path, config),
 
 
 
 
 
 
 
294
  "retrieve_doc": (token_count <= 8000),
295
  "source": config["source_name"],
296
  "content": cleaned_content,
 
18
  Usage:
19
  python process_md_files.py <source1> <source2> ...
20
 
21
+ Where <source1>, <source2>, etc. are one or more of the active sources in
22
+ source_registry.py
23
  (e.g., 'transformers', 'llama_index', 'openai_cookbooks').
24
 
25
  The script processes all Markdown files in the specified input directories (and their subdirectories),
26
  applies the configured filters, and saves the results in JSONL files. Each line in the output
27
  files represents a single document with metadata and content.
28
 
29
+ To add, modify, or retire sources, update data/scraping_scripts/source_registry.py.
30
  """
31
 
32
  import argparse
 
39
 
40
  import tiktoken
41
 
42
+ try:
43
+ from data.scraping_scripts.source_registry import SOURCE_CONFIGS
44
+ except ModuleNotFoundError:
45
+ from source_registry import SOURCE_CONFIGS
46
+
47
  logging.basicConfig(level=logging.INFO)
48
  logger = logging.getLogger(__name__)
49
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
50
 
51
+ def split_frontmatter(content: str) -> tuple[str | None, str]:
52
+ lines = content.splitlines()
53
+ if not lines or lines[0].strip() != "---":
54
+ return None, content
55
 
56
+ for index, line in enumerate(lines[1:], start=1):
57
+ if line.strip() in {"---", "..."}:
58
+ frontmatter = "\n".join(lines[1:index])
59
+ body = "\n".join(lines[index + 1 :])
60
+ return frontmatter, body
61
+
62
+ return None, content
63
+
64
+
65
+ def clean_frontmatter_scalar(value: str) -> str | None:
66
+ value = value.strip()
67
+ if not value or value in {"|", ">"}:
68
+ return None
69
+
70
+ if len(value) >= 2 and value[0] == value[-1] and value[0] in {"'", '"'}:
71
+ quote = value[0]
72
+ value = value[1:-1]
73
+ if quote == '"':
74
+ value = value.replace(r"\"", '"').replace(r"\\", "\\")
75
+ else:
76
+ value = value.replace("''", "'")
77
+
78
+ value = value.strip()
79
+ return value or None
80
+
81
+
82
+ def extract_frontmatter_title(frontmatter: str) -> str | None:
83
+ for key in ("title", "sidebarTitle"):
84
+ title_match = re.search(rf"(?m)^\s*{key}\s*:\s*(.+?)\s*$", frontmatter)
85
+ if title_match:
86
+ title = clean_frontmatter_scalar(title_match.group(1))
87
+ if title:
88
+ return title
89
+
90
+ return None
91
+
92
+
93
+ def iter_markdown_lines_outside_code_fences(content: str):
94
+ in_code_fence = False
95
+ fence_char = None
96
+
97
+ for line in content.splitlines():
98
+ fence_match = re.match(r"^\s{0,3}(```+|~~~+)", line)
99
+ if fence_match:
100
+ current_fence_char = fence_match.group(1)[0]
101
+ if not in_code_fence:
102
+ in_code_fence = True
103
+ fence_char = current_fence_char
104
+ elif current_fence_char == fence_char:
105
+ in_code_fence = False
106
+ fence_char = None
107
+ continue
108
+
109
+ if not in_code_fence:
110
+ yield line
111
 
112
+
113
+ def extract_title(content: str):
114
+ frontmatter, body = split_frontmatter(content)
115
+ if frontmatter:
116
+ title = extract_frontmatter_title(frontmatter)
117
+ if title:
118
+ return title
119
+
120
+ for line in iter_markdown_lines_outside_code_fences(body):
121
+ title_match = re.match(r"^\s{0,3}#\s+(.+)$", line)
122
+ if title_match:
123
+ return title_match.group(1).strip()
124
+
125
+ lines = body.split("\n")
126
  for line in lines:
127
  if line.strip():
128
  return line.strip()
 
130
  return None
131
 
132
 
133
+ def load_source_extension_manifest(directory: str) -> Dict[str, str]:
134
+ manifest_path = os.path.join(directory, "_source_extensions.json")
135
+ if not os.path.exists(manifest_path):
136
+ return {}
137
+
138
+ try:
139
+ with open(manifest_path, "r", encoding="utf-8") as f:
140
+ manifest = json.load(f)
141
+ except (OSError, json.JSONDecodeError):
142
+ logger.warning("Could not load source extension manifest: %s", manifest_path)
143
+ return {}
144
+
145
+ return {
146
+ str(path): str(extension)
147
+ for path, extension in manifest.items()
148
+ if isinstance(path, str) and isinstance(extension, str)
149
+ }
150
+
151
+
152
+ def load_source_url_manifest(directory: str) -> Dict[str, str]:
153
+ manifest_path = os.path.join(directory, "_source_urls.json")
154
+ if not os.path.exists(manifest_path):
155
+ return {}
156
+
157
+ try:
158
+ with open(manifest_path, "r", encoding="utf-8") as f:
159
+ manifest = json.load(f)
160
+ except (OSError, json.JSONDecodeError):
161
+ logger.warning("Could not load source URL manifest: %s", manifest_path)
162
+ return {}
163
+
164
+ return {
165
+ str(path): str(url)
166
+ for path, url in manifest.items()
167
+ if isinstance(path, str) and isinstance(url, str)
168
+ }
169
+
170
+
171
+ def generate_url(
172
+ file_path: str,
173
+ config: Dict,
174
+ source_extension: str | None = None,
175
+ source_url: str | None = None,
176
+ ) -> str:
177
  """
178
  Return an empty string if base_url is empty;
179
  otherwise return the constructed URL as before.
180
  """
181
+ if source_url:
182
+ return source_url
183
+
184
  if not config["base_url"]:
185
  return ""
186
 
 
187
  source_name = config["source_name"]
188
+ path_with_forward_slashes = file_path.replace("\\", "/")
189
+
190
+ if config.get("preserve_file_extension_in_url"):
191
+ if source_extension:
192
+ path_without_extension = os.path.splitext(path_with_forward_slashes)[0]
193
+ return config["base_url"] + path_without_extension + source_extension
194
+ return config["base_url"] + path_with_forward_slashes
195
+
196
+ path_without_extension = os.path.splitext(file_path)[0]
197
 
198
  if source_name == "llama_index":
199
  framework_prefix = "src/content/docs/framework/"
 
236
  return cleaned_content.strip()
237
 
238
 
239
+ def remove_inline_base64_images(content: str) -> str:
240
+ content = re.sub(
241
+ r"!\[[^\]]*\]\(\s*data:image/[^,\s)]+;base64,[A-Za-z0-9+/=]+(?:\s+\"[^\"]*\")?\s*\)",
242
+ "[inline image omitted]",
243
+ content,
244
+ )
245
+ content = re.sub(
246
+ r"<img\b[^>]*\bsrc=[\"']data:image/[^,\s\"']+;base64,[A-Za-z0-9+/=]+[\"'][^>]*>",
247
+ "[inline image omitted]",
248
+ content,
249
+ flags=re.IGNORECASE,
250
+ )
251
+ content = re.sub(
252
+ r"data:image/[^,\s)\"']+;base64,[A-Za-z0-9+/=]+",
253
+ "[inline image omitted]",
254
+ content,
255
+ )
256
+ return content
257
+
258
+
259
+ def clean_document_content(content: str) -> str:
260
+ content = remove_copyright_header(content)
261
+ content = remove_inline_base64_images(content)
262
+ return content.strip()
263
+
264
+
265
  def process_md_files(directory: str, config: Dict) -> List[Dict]:
266
  jsonl_data = []
267
+ source_extension_manifest = load_source_extension_manifest(directory)
268
+ source_url_manifest = load_source_url_manifest(directory)
269
 
270
  for root, _, files in os.walk(directory):
271
  for file in files:
 
277
  with open(file_path, "r", encoding="utf-8") as f:
278
  content = f.read()
279
 
280
+ cleaned_content = clean_document_content(content)
281
+ title = extract_title(cleaned_content)
282
+ token_count = num_tokens_from_string(cleaned_content, "cl100k_base")
283
 
284
  # Skip very small or extremely large files
285
  if token_count < 100 or token_count > 200_000:
 
288
  )
289
  continue
290
 
 
 
291
  json_object = {
292
  "tokens": token_count,
293
  "doc_id": str(uuid.uuid4()),
294
  "name": (title if title else file),
295
+ "url": generate_url(
296
+ relative_path,
297
+ config,
298
+ source_extension_manifest.get(
299
+ relative_path.replace("\\", "/")
300
+ ),
301
+ source_url_manifest.get(relative_path.replace("\\", "/")),
302
+ ),
303
  "retrieve_doc": (token_count <= 8000),
304
  "source": config["source_name"],
305
  "content": cleaned_content,
data/scraping_scripts/retire_source_workflow.py CHANGED
@@ -5,8 +5,9 @@ Retire one or more sources from the AI Tutor data pipeline.
5
  This removes the source from:
6
  1. data/all_sources_data.jsonl
7
  2. data/all_sources_contextual_nodes.pkl
8
- 3. the rebuilt Chroma vector store
9
- 4. the private Hugging Face data repository's per-source JSONL file
 
10
 
11
  Example:
12
  uv run -m data.scraping_scripts.retire_source_workflow --sources 8-hour_primer --yes
@@ -15,9 +16,11 @@ Example:
15
  from __future__ import annotations
16
 
17
  import argparse
 
18
  import json
19
  import os
20
  import pickle
 
21
  import shutil
22
  import subprocess
23
  import sys
@@ -33,7 +36,7 @@ from data.scraping_scripts.hf_auth import HuggingFaceAuthError, validate_hf_acce
33
  from scripts.chroma_rag import get_chunk_record_source
34
 
35
  try:
36
- from data.scraping_scripts.process_md_files import SOURCE_CONFIGS
37
  except Exception:
38
  SOURCE_CONFIGS = {}
39
 
@@ -43,6 +46,7 @@ DATA_REPO_ID = "towardsai-tutors/ai-tutor-data"
43
  VECTOR_REPO_ID = "towardsai-tutors/ai-tutor-vector-db"
44
  ALL_SOURCES_JSONL = Path("data/all_sources_data.jsonl")
45
  CONTEXTUAL_NODES = Path("data/all_sources_contextual_nodes.pkl")
 
46
 
47
 
48
  def run_module(module_name: str, *module_args: str) -> subprocess.CompletedProcess:
@@ -168,6 +172,104 @@ def delete_local_source_files(sources: list[str], dry_run: bool) -> None:
168
  print(f"Deleted local source file {source_path} after backup to {backup_path}")
169
 
170
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
171
  def rebuild_vector_store() -> None:
172
  if not os.getenv("COHERE_API_KEY"):
173
  raise SystemExit("COHERE_API_KEY is required to rebuild the Chroma vector store.")
@@ -246,7 +348,7 @@ def main() -> None:
246
  nargs="*",
247
  help=(
248
  "Optional exact per-source JSONL filenames to delete from the data repo. "
249
- "Defaults to SOURCE_CONFIGS output files or <source>_data.jsonl."
250
  ),
251
  )
252
  parser.add_argument(
@@ -279,6 +381,11 @@ def main() -> None:
279
  action="store_true",
280
  help="Delete matching local per-source JSONL files after backing them up.",
281
  )
 
 
 
 
 
282
  parser.add_argument(
283
  "--dry-run",
284
  action="store_true",
@@ -333,6 +440,9 @@ def main() -> None:
333
  if args.delete_local_source_files:
334
  delete_local_source_files(args.sources, args.dry_run)
335
 
 
 
 
336
  if args.dry_run:
337
  print("Dry run complete. No files were changed.")
338
  return
 
5
  This removes the source from:
6
  1. data/all_sources_data.jsonl
7
  2. data/all_sources_contextual_nodes.pkl
8
+ 3. data/scraping_scripts/source_registry.py
9
+ 4. the rebuilt Chroma vector store
10
+ 5. the private Hugging Face data repository's per-source JSONL file
11
 
12
  Example:
13
  uv run -m data.scraping_scripts.retire_source_workflow --sources 8-hour_primer --yes
 
16
  from __future__ import annotations
17
 
18
  import argparse
19
+ import ast
20
  import json
21
  import os
22
  import pickle
23
+ import re
24
  import shutil
25
  import subprocess
26
  import sys
 
36
  from scripts.chroma_rag import get_chunk_record_source
37
 
38
  try:
39
+ from data.scraping_scripts.source_registry import SOURCE_CONFIGS
40
  except Exception:
41
  SOURCE_CONFIGS = {}
42
 
 
46
  VECTOR_REPO_ID = "towardsai-tutors/ai-tutor-vector-db"
47
  ALL_SOURCES_JSONL = Path("data/all_sources_data.jsonl")
48
  CONTEXTUAL_NODES = Path("data/all_sources_contextual_nodes.pkl")
49
+ SOURCE_REGISTRY = Path("data/scraping_scripts/source_registry.py")
50
 
51
 
52
  def run_module(module_name: str, *module_args: str) -> subprocess.CompletedProcess:
 
172
  print(f"Deleted local source file {source_path} after backup to {backup_path}")
173
 
174
 
175
+ def find_source_config_block(text: str, source: str) -> tuple[int, int] | None:
176
+ marker = f' "{source}": {{'
177
+ start = text.find(marker)
178
+ if start == -1:
179
+ return None
180
+
181
+ body_start = text.find("{", start)
182
+ if body_start == -1:
183
+ return None
184
+
185
+ depth = 0
186
+ quote: str | None = None
187
+ escaped = False
188
+ for index in range(body_start, len(text)):
189
+ char = text[index]
190
+ if quote:
191
+ if escaped:
192
+ escaped = False
193
+ elif char == "\\":
194
+ escaped = True
195
+ elif char == quote:
196
+ quote = None
197
+ continue
198
+
199
+ if char in {"'", '"'}:
200
+ quote = char
201
+ elif char == "{":
202
+ depth += 1
203
+ elif char == "}":
204
+ depth -= 1
205
+ if depth == 0:
206
+ end = index + 1
207
+ if end < len(text) and text[end] == ",":
208
+ end += 1
209
+ if end < len(text) and text[end] == "\n":
210
+ end += 1
211
+ return start, end
212
+
213
+ return None
214
+
215
+
216
+ def remove_source_from_registry_text(text: str, source: str) -> tuple[str, bool]:
217
+ original = text
218
+ block = find_source_config_block(text, source)
219
+ if block is not None:
220
+ start, end = block
221
+ text = text[:start] + text[end:]
222
+
223
+ source_pattern = re.escape(source)
224
+ text = re.sub(
225
+ rf'^[ \t]*"{source_pattern}",\n',
226
+ "",
227
+ text,
228
+ flags=re.MULTILINE,
229
+ )
230
+ text = re.sub(
231
+ rf'^[ \t]*"{source_pattern}":\s*"[^"\n]*",\n',
232
+ "",
233
+ text,
234
+ flags=re.MULTILINE,
235
+ )
236
+ return text, text != original
237
+
238
+
239
+ def update_source_registry(sources: list[str], dry_run: bool) -> None:
240
+ if not SOURCE_REGISTRY.exists():
241
+ print(f"Warning: {SOURCE_REGISTRY} not found; registry was not updated.")
242
+ return
243
+
244
+ original = SOURCE_REGISTRY.read_text(encoding="utf-8")
245
+ updated = original
246
+ changed_sources: list[str] = []
247
+
248
+ for source in sources:
249
+ updated, changed = remove_source_from_registry_text(updated, source)
250
+ if changed:
251
+ changed_sources.append(source)
252
+
253
+ if not changed_sources:
254
+ print(f"No retired sources were found in {SOURCE_REGISTRY}.")
255
+ return
256
+
257
+ if dry_run:
258
+ print(
259
+ "Would remove from source_registry.py: "
260
+ + ", ".join(sorted(changed_sources))
261
+ )
262
+ return
263
+
264
+ ast.parse(updated, filename=str(SOURCE_REGISTRY))
265
+ backup_path = backup_file(SOURCE_REGISTRY)
266
+ SOURCE_REGISTRY.write_text(updated, encoding="utf-8")
267
+ print(
268
+ f"Removed {', '.join(sorted(changed_sources))} from {SOURCE_REGISTRY} "
269
+ f"after backup to {backup_path}"
270
+ )
271
+
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.")
 
348
  nargs="*",
349
  help=(
350
  "Optional exact per-source JSONL filenames to delete from the data repo. "
351
+ "Defaults to source_registry.py output files or <source>_data.jsonl."
352
  ),
353
  )
354
  parser.add_argument(
 
381
  action="store_true",
382
  help="Delete matching local per-source JSONL files after backing them up.",
383
  )
384
+ parser.add_argument(
385
+ "--keep-source-registry",
386
+ action="store_true",
387
+ help="Do not remove retired sources from data/scraping_scripts/source_registry.py.",
388
+ )
389
  parser.add_argument(
390
  "--dry-run",
391
  action="store_true",
 
440
  if args.delete_local_source_files:
441
  delete_local_source_files(args.sources, args.dry_run)
442
 
443
+ if not args.keep_source_registry:
444
+ update_source_registry(args.sources, args.dry_run)
445
+
446
  if args.dry_run:
447
  print("Dry run complete. No files were changed.")
448
  return
data/scraping_scripts/source_registry.py ADDED
@@ -0,0 +1,358 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Central source registry for the AI Tutor knowledge base.
2
+
3
+ Sources listed here are active in the KB pipeline. To retire a source, run
4
+ ``retire_source_workflow.py``; confirmed retirements remove sources from this
5
+ file automatically.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ from pathlib import Path
11
+ from typing import Any, Iterable
12
+
13
+ ALL_SOURCES_JSONL = "data/all_sources_data.jsonl"
14
+ CONTEXTUAL_NODES_PKL = "data/all_sources_contextual_nodes.pkl"
15
+
16
+
17
+ SOURCE_CONFIGS: dict[str, dict[str, Any]] = {
18
+ "transformers": {
19
+ "base_url": "https://huggingface.co/docs/transformers/",
20
+ "input_directory": "data/transformers_md_files",
21
+ "output_file": "data/transformers_data.jsonl",
22
+ "source_name": "transformers",
23
+ "use_include_list": False,
24
+ "included_dirs": [],
25
+ "excluded_dirs": ["internal"],
26
+ "excluded_root_files": [],
27
+ "included_root_files": [],
28
+ "url_extension": "",
29
+ },
30
+ "peft": {
31
+ "base_url": "https://huggingface.co/docs/peft/",
32
+ "input_directory": "data/peft_md_files",
33
+ "output_file": "data/peft_data.jsonl",
34
+ "source_name": "peft",
35
+ "use_include_list": False,
36
+ "included_dirs": [],
37
+ "excluded_dirs": [],
38
+ "excluded_root_files": [],
39
+ "included_root_files": [],
40
+ "url_extension": "",
41
+ },
42
+ "trl": {
43
+ "base_url": "https://huggingface.co/docs/trl/",
44
+ "input_directory": "data/trl_md_files",
45
+ "output_file": "data/trl_data.jsonl",
46
+ "source_name": "trl",
47
+ "use_include_list": False,
48
+ "included_dirs": [],
49
+ "excluded_dirs": [],
50
+ "excluded_root_files": [],
51
+ "included_root_files": [],
52
+ "url_extension": "",
53
+ },
54
+ "llama_index": {
55
+ "base_url": "https://docs.llamaindex.ai/en/stable/",
56
+ "input_directory": "data/llama_index_md_files",
57
+ "output_file": "data/llama_index_data.jsonl",
58
+ "source_name": "llama_index",
59
+ "use_include_list": True,
60
+ "included_dirs": [
61
+ "src/content/docs/framework/index.md",
62
+ "src/content/docs/framework/getting_started",
63
+ "src/content/docs/framework/understanding",
64
+ "src/content/docs/framework/use_cases",
65
+ "src/content/docs/framework/module_guides",
66
+ "src/content/docs/framework/optimizing",
67
+ "src/content/docs/framework/community/faq",
68
+ "src/content/docs/framework/community/integrations",
69
+ "src/content/docs/framework/llama_cloud",
70
+ "examples",
71
+ ],
72
+ "excluded_dirs": [],
73
+ "excluded_root_files": [],
74
+ "included_root_files": [],
75
+ "url_extension": "",
76
+ },
77
+ "langchain": {
78
+ "base_url": "https://docs.langchain.com/oss/python/",
79
+ "input_directory": "data/langchain_md_files",
80
+ "output_file": "data/langchain_data.jsonl",
81
+ "source_name": "langchain",
82
+ "use_include_list": True,
83
+ "included_dirs": [
84
+ "concepts",
85
+ "langchain",
86
+ "python/integrations/chat/",
87
+ "python/integrations/document_loaders/",
88
+ "python/integrations/document_transformers/",
89
+ "python/integrations/embeddings/",
90
+ "python/integrations/retrievers/",
91
+ "python/integrations/splitters/",
92
+ "python/integrations/stores/",
93
+ "python/integrations/tools/",
94
+ "python/integrations/vectorstores/",
95
+ "python/migrate",
96
+ "python/releases",
97
+ ],
98
+ "excluded_dirs": [],
99
+ "excluded_root_files": [],
100
+ "included_root_files": [
101
+ "security-policy.mdx",
102
+ "release-policy.mdx",
103
+ "versioning.mdx",
104
+ ],
105
+ "url_extension": "",
106
+ },
107
+ "langgraph": {
108
+ "base_url": "https://docs.langchain.com/oss/python/langgraph/",
109
+ "input_directory": "data/langgraph_md_files",
110
+ "output_file": "data/langgraph_data.jsonl",
111
+ "source_name": "langgraph",
112
+ "use_include_list": False,
113
+ "included_dirs": [],
114
+ "excluded_dirs": [],
115
+ "excluded_root_files": [],
116
+ "included_root_files": [],
117
+ "url_extension": "",
118
+ },
119
+ "deep_agents": {
120
+ "base_url": "https://docs.langchain.com/oss/python/deepagents/",
121
+ "input_directory": "data/deep_agents_md_files",
122
+ "output_file": "data/deep_agents_data.jsonl",
123
+ "source_name": "deep_agents",
124
+ "use_include_list": False,
125
+ "included_dirs": [],
126
+ "excluded_dirs": [],
127
+ "excluded_root_files": [],
128
+ "included_root_files": [],
129
+ "url_extension": "",
130
+ },
131
+ "openai_docs": {
132
+ "base_url": "https://developers.openai.com/",
133
+ "input_directory": "data/openai_docs_md_files",
134
+ "output_file": "data/openai_docs_data.jsonl",
135
+ "source_name": "openai_docs",
136
+ "use_include_list": False,
137
+ "included_dirs": [],
138
+ "excluded_dirs": [],
139
+ "excluded_root_files": [],
140
+ "included_root_files": [],
141
+ "url_extension": "",
142
+ "llms_txt_urls": [
143
+ "https://developers.openai.com/api/docs/llms.txt",
144
+ "https://developers.openai.com/codex/llms.txt",
145
+ ],
146
+ "llms_url_include_prefixes": [
147
+ "https://developers.openai.com/api/docs/",
148
+ "https://developers.openai.com/codex/",
149
+ ],
150
+ },
151
+ "claude_code_docs": {
152
+ "base_url": "https://code.claude.com/docs/",
153
+ "input_directory": "data/claude_code_docs_md_files",
154
+ "output_file": "data/claude_code_docs_data.jsonl",
155
+ "source_name": "claude_code_docs",
156
+ "use_include_list": False,
157
+ "included_dirs": [],
158
+ "excluded_dirs": [],
159
+ "excluded_root_files": [],
160
+ "included_root_files": [],
161
+ "url_extension": "",
162
+ "llms_txt_urls": [
163
+ "https://code.claude.com/docs/llms.txt",
164
+ ],
165
+ "llms_url_include_prefixes": [
166
+ "https://code.claude.com/docs/en/",
167
+ ],
168
+ },
169
+ "full_stack_ai_engineering": {
170
+ "base_url": "",
171
+ "input_directory": "data/full_stack_ai_engineering",
172
+ "output_file": "data/full_stack_ai_engineering_data.jsonl",
173
+ "source_name": "full_stack_ai_engineering",
174
+ "use_include_list": False,
175
+ "included_dirs": [],
176
+ "excluded_dirs": [],
177
+ "excluded_root_files": [],
178
+ "included_root_files": [],
179
+ "url_extension": "",
180
+ },
181
+ "beginner_python_for_ai_engineering": {
182
+ "base_url": "",
183
+ "input_directory": "data/beginner_python_for_ai_engineering",
184
+ "output_file": "data/beginner_python_for_ai_engineering_data.jsonl",
185
+ "source_name": "beginner_python_for_ai_engineering",
186
+ "use_include_list": False,
187
+ "included_dirs": [],
188
+ "excluded_dirs": [],
189
+ "excluded_root_files": [],
190
+ "included_root_files": [],
191
+ "url_extension": "",
192
+ },
193
+ "master_ai_for_work": {
194
+ "base_url": "",
195
+ "input_directory": "data/master_ai_for_work",
196
+ "output_file": "data/master_ai_for_work_data.jsonl",
197
+ "source_name": "master_ai_for_work",
198
+ "use_include_list": False,
199
+ "included_dirs": [],
200
+ "excluded_dirs": [],
201
+ "excluded_root_files": [],
202
+ "included_root_files": [],
203
+ "url_extension": "",
204
+ },
205
+ "agentic_ai_engineering": {
206
+ "base_url": "",
207
+ "input_directory": "data/agentic_ai_engineering",
208
+ "output_file": "data/agentic_ai_engineering_data.jsonl",
209
+ "source_name": "agentic_ai_engineering",
210
+ "use_include_list": False,
211
+ "included_dirs": [],
212
+ "excluded_dirs": [],
213
+ "excluded_root_files": [],
214
+ "included_root_files": [],
215
+ "url_extension": "",
216
+ },
217
+ }
218
+
219
+ DOC_SOURCE_KEYS = (
220
+ "transformers",
221
+ "peft",
222
+ "trl",
223
+ "llama_index",
224
+ "langchain",
225
+ "langgraph",
226
+ "deep_agents",
227
+ "openai_docs",
228
+ "claude_code_docs",
229
+ )
230
+ GITHUB_SOURCE_KEYS = (
231
+ "transformers",
232
+ "peft",
233
+ "trl",
234
+ "llama_index",
235
+ "langchain",
236
+ "langgraph",
237
+ "deep_agents",
238
+ )
239
+ LLMS_TXT_SOURCE_KEYS = (
240
+ "openai_docs",
241
+ "claude_code_docs",
242
+ )
243
+ COURSE_SOURCE_KEYS = frozenset(
244
+ {
245
+ "full_stack_ai_engineering",
246
+ "beginner_python_for_ai_engineering",
247
+ "master_ai_for_work",
248
+ "agentic_ai_engineering",
249
+ }
250
+ )
251
+
252
+ ACTIVE_SOURCE_KEYS = frozenset(SOURCE_CONFIGS.keys())
253
+ AVAILABLE_SOURCES = list(SOURCE_CONFIGS.keys())
254
+
255
+ SOURCE_KEY_TO_LABEL = {
256
+ "transformers": "Transformers Docs",
257
+ "peft": "PEFT Docs",
258
+ "trl": "TRL Docs",
259
+ "llama_index": "LlamaIndex Docs",
260
+ "langchain": "LangChain Docs",
261
+ "langgraph": "LangGraph Docs",
262
+ "deep_agents": "Deep Agents Docs",
263
+ "openai_docs": "OpenAI Docs",
264
+ "claude_code_docs": "Claude Code Docs",
265
+ "full_stack_ai_engineering": "Full Stack AI Engineering",
266
+ "beginner_python_for_ai_engineering": "Beginner Python for AI Engineering",
267
+ "master_ai_for_work": "Master AI For Work",
268
+ "agentic_ai_engineering": "Agentic AI Engineering",
269
+ }
270
+
271
+ UI_SOURCE_KEYS = (
272
+ "openai_docs",
273
+ "claude_code_docs",
274
+ "langgraph",
275
+ "deep_agents",
276
+ "langchain",
277
+ "llama_index",
278
+ "transformers",
279
+ "peft",
280
+ "trl",
281
+ "agentic_ai_engineering",
282
+ "master_ai_for_work",
283
+ "full_stack_ai_engineering",
284
+ "beginner_python_for_ai_engineering",
285
+ )
286
+ SOURCE_UI_TO_KEY = {SOURCE_KEY_TO_LABEL[key]: key for key in UI_SOURCE_KEYS}
287
+ AVAILABLE_SOURCES_UI = list(SOURCE_UI_TO_KEY.keys())
288
+
289
+ DEFAULT_SELECTED_SOURCE_KEYS = (
290
+ "agentic_ai_engineering",
291
+ "master_ai_for_work",
292
+ "full_stack_ai_engineering",
293
+ "beginner_python_for_ai_engineering",
294
+ "openai_docs",
295
+ "claude_code_docs",
296
+ "langgraph",
297
+ "deep_agents",
298
+ "transformers",
299
+ "peft",
300
+ "trl",
301
+ "llama_index",
302
+ "langchain",
303
+ )
304
+ DEFAULT_SELECTED_SOURCES_UI = [
305
+ SOURCE_KEY_TO_LABEL[key] for key in DEFAULT_SELECTED_SOURCE_KEYS
306
+ ]
307
+
308
+
309
+ def aggregate_data_files() -> dict[str, str]:
310
+ return {
311
+ ALL_SOURCES_JSONL: Path(ALL_SOURCES_JSONL).name,
312
+ CONTEXTUAL_NODES_PKL: Path(CONTEXTUAL_NODES_PKL).name,
313
+ }
314
+
315
+
316
+ def source_data_files() -> dict[str, str]:
317
+ return {
318
+ str(config["output_file"]): Path(str(config["output_file"])).name
319
+ for config in SOURCE_CONFIGS.values()
320
+ }
321
+
322
+
323
+ def source_output_files(sources: Iterable[str]) -> set[str]:
324
+ return {
325
+ str(SOURCE_CONFIGS[source]["output_file"])
326
+ for source in sources
327
+ if source in SOURCE_CONFIGS
328
+ }
329
+
330
+
331
+ def required_data_files(*, include_source_files: bool = True) -> dict[str, str]:
332
+ files = aggregate_data_files()
333
+ if include_source_files:
334
+ files.update(source_data_files())
335
+ return files
336
+
337
+
338
+ def upload_data_file_paths() -> list[str]:
339
+ return list(required_data_files().keys())
340
+
341
+
342
+ def vector_store_source_configs() -> dict[str, dict[str, str]]:
343
+ configs = {
344
+ source: {
345
+ "input_file": str(config["output_file"]),
346
+ "db_name": f"chroma-db-{source}",
347
+ "document_dict_file": f"document_dict_{source}.pkl",
348
+ "bm25_index_file": f"bm25_index_{source}.pkl",
349
+ }
350
+ for source, config in SOURCE_CONFIGS.items()
351
+ }
352
+ configs["all_sources"] = {
353
+ "input_file": ALL_SOURCES_JSONL,
354
+ "db_name": "chroma-db-all_sources",
355
+ "document_dict_file": "document_dict_all_sources.pkl",
356
+ "bm25_index_file": "bm25_index_all_sources.pkl",
357
+ }
358
+ return configs
data/scraping_scripts/update_docs_workflow.py CHANGED
@@ -2,14 +2,14 @@
2
  """
3
  AI Tutor App - Documentation Update Workflow
4
 
5
- This script automates the process of updating documentation from GitHub repositories:
6
- 1. Download documentation from GitHub using the API
7
  2. Process markdown files to create JSONL data
8
  3. Add contextual information to document nodes
9
  4. Create vector stores
10
  5. Upload databases to HuggingFace
11
 
12
- This workflow is specific to updating library documentation (Transformers, PEFT, LlamaIndex, etc.).
13
  For adding courses, use the add_course_workflow.py script instead.
14
 
15
  Usage:
@@ -36,7 +36,17 @@ from typing import Dict, List, Set
36
  from dotenv import load_dotenv
37
  from huggingface_hub import hf_hub_download
38
 
 
 
 
39
  from data.scraping_scripts.hf_auth import HuggingFaceAuthError, validate_hf_access
 
 
 
 
 
 
 
40
  from scripts.chroma_rag import get_chunk_record_doc_id
41
 
42
  # Load environment variables from .env file
@@ -62,27 +72,10 @@ def ensure_hf_access() -> None:
62
  sys.exit(1)
63
 
64
 
65
- def ensure_required_files_exist():
66
  """Download required data files from HuggingFace if they don't exist locally."""
67
- # List of files to check and download
68
- required_files = {
69
- # Critical files
70
- "data/all_sources_data.jsonl": "all_sources_data.jsonl",
71
- "data/all_sources_contextual_nodes.pkl": "all_sources_contextual_nodes.pkl",
72
- # Documentation source files
73
- "data/transformers_data.jsonl": "transformers_data.jsonl",
74
- "data/peft_data.jsonl": "peft_data.jsonl",
75
- "data/trl_data.jsonl": "trl_data.jsonl",
76
- "data/llama_index_data.jsonl": "llama_index_data.jsonl",
77
- "data/langchain_data.jsonl": "langchain_data.jsonl",
78
- "data/openai_cookbooks_data.jsonl": "openai_cookbooks_data.jsonl",
79
- # Course files
80
- "data/tai_blog_data.jsonl": "tai_blog_data.jsonl",
81
- "data/master_ai_for_work_data.jsonl": "master_ai_for_work_data.jsonl",
82
- "data/agentic_ai_engineering_data.jsonl": "agentic_ai_engineering_data.jsonl",
83
- "data/full_stack_ai_engineering_data.jsonl": "full_stack_ai_engineering_data.jsonl",
84
- "data/beginner_python_for_ai_engineering_data.jsonl": "beginner_python_for_ai_engineering_data.jsonl",
85
- }
86
 
87
  # Critical files that must be downloaded
88
  critical_files = [
@@ -92,6 +85,14 @@ def ensure_required_files_exist():
92
 
93
  # Check and download each file
94
  for local_path, remote_filename in required_files.items():
 
 
 
 
 
 
 
 
95
  if not os.path.exists(local_path):
96
  logger.info(
97
  f"{remote_filename} not found. Attempting to download from HuggingFace..."
@@ -129,15 +130,10 @@ def ensure_required_files_exist():
129
  )
130
 
131
 
132
- # Documentation sources that can be updated via GitHub API
133
- GITHUB_SOURCES = [
134
- "transformers",
135
- "peft",
136
- "trl",
137
- "llama_index",
138
- "openai_cookbooks",
139
- "langchain",
140
- ]
141
 
142
 
143
  def load_jsonl(file_path: str) -> List[Dict]:
@@ -178,6 +174,46 @@ def download_from_github(sources: List[str]) -> None:
178
  logger.info(f"Successfully downloaded {source} documentation")
179
 
180
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
181
  def capture_source_versions(sources: List[str]) -> None:
182
  """Record latest release tag + SHA + indexed date per source."""
183
  logger.info(f"Capturing source versions for: {sources}")
@@ -202,6 +238,12 @@ def process_markdown_files(sources: List[str]) -> None:
202
  logger.error("Error processing markdown files - check output above")
203
  sys.exit(1)
204
 
 
 
 
 
 
 
205
  logger.info("Successfully processed markdown files")
206
 
207
 
@@ -367,12 +409,12 @@ def main():
367
  parser.add_argument(
368
  "--sources",
369
  nargs="+",
370
- choices=GITHUB_SOURCES,
371
- default=GITHUB_SOURCES,
372
- help="GitHub documentation sources to update",
373
  )
374
  parser.add_argument(
375
- "--skip-download", action="store_true", help="Skip downloading from GitHub"
376
  )
377
  parser.add_argument(
378
  "--skip-process", action="store_true", help="Skip processing markdown files"
@@ -403,12 +445,14 @@ def main():
403
 
404
  ensure_hf_access()
405
 
406
- # Ensure required data files exist before proceeding
407
- ensure_required_files_exist()
 
 
408
 
409
  # Execute the workflow steps
410
  if not args.skip_download:
411
- download_from_github(args.sources)
412
  capture_source_versions(args.sources)
413
 
414
  if not args.skip_process:
@@ -417,6 +461,8 @@ def main():
417
  if not args.skip_context:
418
  add_context_to_nodes(not args.process_all_context)
419
 
 
 
420
  if not args.skip_vectors:
421
  create_vector_stores()
422
 
 
2
  """
3
  AI Tutor App - Documentation Update Workflow
4
 
5
+ This script automates the process of updating documentation sources:
6
+ 1. Download documentation from GitHub or official llms.txt indexes
7
  2. Process markdown files to create JSONL data
8
  3. Add contextual information to document nodes
9
  4. Create vector stores
10
  5. Upload databases to HuggingFace
11
 
12
+ This workflow is specific to updating library documentation (Transformers, PEFT, LlamaIndex, OpenAI docs, etc.).
13
  For adding courses, use the add_course_workflow.py script instead.
14
 
15
  Usage:
 
36
  from dotenv import load_dotenv
37
  from huggingface_hub import hf_hub_download
38
 
39
+ from data.scraping_scripts.contextual_node_pruning import (
40
+ prune_contextual_nodes_to_active_sources,
41
+ )
42
  from data.scraping_scripts.hf_auth import HuggingFaceAuthError, validate_hf_access
43
+ from data.scraping_scripts.source_registry import (
44
+ DOC_SOURCE_KEYS,
45
+ GITHUB_SOURCE_KEYS,
46
+ LLMS_TXT_SOURCE_KEYS,
47
+ required_data_files,
48
+ source_output_files,
49
+ )
50
  from scripts.chroma_rag import get_chunk_record_doc_id
51
 
52
  # Load environment variables from .env file
 
72
  sys.exit(1)
73
 
74
 
75
+ def ensure_required_files_exist(sources_to_regenerate: List[str] | None = None):
76
  """Download required data files from HuggingFace if they don't exist locally."""
77
+ required_files = required_data_files()
78
+ regenerated_source_files = source_output_files(sources_to_regenerate or [])
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
79
 
80
  # Critical files that must be downloaded
81
  critical_files = [
 
85
 
86
  # Check and download each file
87
  for local_path, remote_filename in required_files.items():
88
+ if local_path in regenerated_source_files:
89
+ if not os.path.exists(local_path):
90
+ logger.info(
91
+ "%s will be regenerated for this run; skipping HuggingFace download",
92
+ remote_filename,
93
+ )
94
+ continue
95
+
96
  if not os.path.exists(local_path):
97
  logger.info(
98
  f"{remote_filename} not found. Attempting to download from HuggingFace..."
 
130
  )
131
 
132
 
133
+ # Documentation sources that can be updated automatically
134
+ DOC_SOURCES = list(DOC_SOURCE_KEYS)
135
+ GITHUB_SOURCES = list(GITHUB_SOURCE_KEYS)
136
+ LLMS_TXT_SOURCES = list(LLMS_TXT_SOURCE_KEYS)
 
 
 
 
 
137
 
138
 
139
  def load_jsonl(file_path: str) -> List[Dict]:
 
174
  logger.info(f"Successfully downloaded {source} documentation")
175
 
176
 
177
+ def download_from_llms_txt(sources: List[str]) -> None:
178
+ """Download documentation from llms.txt indexes."""
179
+ logger.info(
180
+ f"Downloading documentation from llms.txt indexes for sources: {sources}"
181
+ )
182
+
183
+ for source in sources:
184
+ if source not in LLMS_TXT_SOURCES:
185
+ logger.warning(
186
+ f"Source {source} is not an llms.txt source, skipping download"
187
+ )
188
+ continue
189
+
190
+ logger.info(f"Downloading {source} documentation")
191
+ result = run_module("data.scraping_scripts.llms_txt_to_markdown_docs", source)
192
+
193
+ if result.returncode != 0:
194
+ logger.error(
195
+ f"Error downloading {source} documentation. Stopping workflow to avoid overwriting source JSONL files with incomplete data."
196
+ )
197
+ sys.exit(1)
198
+
199
+ logger.info(f"Successfully downloaded {source} documentation")
200
+
201
+
202
+ def download_documentation(sources: List[str]) -> None:
203
+ """Download docs with the right source-specific downloader."""
204
+ github_sources = [source for source in sources if source in GITHUB_SOURCES]
205
+ llms_txt_sources = [source for source in sources if source in LLMS_TXT_SOURCES]
206
+
207
+ if github_sources:
208
+ download_from_github(github_sources)
209
+ if llms_txt_sources:
210
+ download_from_llms_txt(llms_txt_sources)
211
+
212
+ unsupported = sorted(set(sources) - set(github_sources) - set(llms_txt_sources))
213
+ for source in unsupported:
214
+ logger.warning(f"Source {source} is not a downloadable docs source, skipping")
215
+
216
+
217
  def capture_source_versions(sources: List[str]) -> None:
218
  """Record latest release tag + SHA + indexed date per source."""
219
  logger.info(f"Capturing source versions for: {sources}")
 
238
  logger.error("Error processing markdown files - check output above")
239
  sys.exit(1)
240
 
241
+ if len(sources) == 1:
242
+ from data.scraping_scripts.process_md_files import combine_all_sources
243
+
244
+ logger.info("Rebuilding all_sources_data.jsonl after single-source update")
245
+ combine_all_sources(sources)
246
+
247
  logger.info("Successfully processed markdown files")
248
 
249
 
 
409
  parser.add_argument(
410
  "--sources",
411
  nargs="+",
412
+ choices=DOC_SOURCES,
413
+ default=DOC_SOURCES,
414
+ help="Documentation sources to update",
415
  )
416
  parser.add_argument(
417
+ "--skip-download", action="store_true", help="Skip downloading source docs"
418
  )
419
  parser.add_argument(
420
  "--skip-process", action="store_true", help="Skip processing markdown files"
 
445
 
446
  ensure_hf_access()
447
 
448
+ # Keep untouched source JSONLs by downloading them when needed, but don't
449
+ # require first-time sources that this run is about to regenerate.
450
+ sources_to_regenerate = [] if args.skip_process else args.sources
451
+ ensure_required_files_exist(sources_to_regenerate=sources_to_regenerate)
452
 
453
  # Execute the workflow steps
454
  if not args.skip_download:
455
+ download_documentation(args.sources)
456
  capture_source_versions(args.sources)
457
 
458
  if not args.skip_process:
 
461
  if not args.skip_context:
462
  add_context_to_nodes(not args.process_all_context)
463
 
464
+ prune_contextual_nodes_to_active_sources()
465
+
466
  if not args.skip_vectors:
467
  create_vector_stores()
468
 
data/scraping_scripts/upload_data_to_hf.py CHANGED
@@ -22,33 +22,17 @@ from dotenv import load_dotenv
22
 
23
  try:
24
  from data.scraping_scripts.hf_auth import HuggingFaceAuthError, validate_hf_access
 
25
  except ModuleNotFoundError:
26
  from hf_auth import HuggingFaceAuthError, validate_hf_access
 
27
 
28
  load_dotenv()
29
 
30
 
31
  def upload_files_to_huggingface(repo_id="towardsai-tutors/ai-tutor-data"):
32
  """Upload data files to a private HuggingFace repository."""
33
- # Main files to upload
34
- files_to_upload = [
35
- # Combined data and vector store
36
- "data/all_sources_data.jsonl",
37
- "data/all_sources_contextual_nodes.pkl",
38
- # Individual source files
39
- "data/transformers_data.jsonl",
40
- "data/peft_data.jsonl",
41
- "data/trl_data.jsonl",
42
- "data/llama_index_data.jsonl",
43
- "data/langchain_data.jsonl",
44
- "data/openai_cookbooks_data.jsonl",
45
- # Course files
46
- "data/tai_blog_data.jsonl",
47
- "data/master_ai_for_work_data.jsonl",
48
- "data/agentic_ai_engineering_data.jsonl",
49
- "data/full_stack_ai_engineering_data.jsonl",
50
- "data/beginner_python_for_ai_engineering_data.jsonl",
51
- ]
52
 
53
  # Filter to only include files that exist
54
  existing_files = []
 
22
 
23
  try:
24
  from data.scraping_scripts.hf_auth import HuggingFaceAuthError, validate_hf_access
25
+ from data.scraping_scripts.source_registry import upload_data_file_paths
26
  except ModuleNotFoundError:
27
  from hf_auth import HuggingFaceAuthError, validate_hf_access
28
+ from source_registry import upload_data_file_paths
29
 
30
  load_dotenv()
31
 
32
 
33
  def upload_files_to_huggingface(repo_id="towardsai-tutors/ai-tutor-data"):
34
  """Upload data files to a private HuggingFace repository."""
35
+ files_to_upload = upload_data_file_paths()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
36
 
37
  # Filter to only include files that exist
38
  existing_files = []
frontend/lib/doc-metadata.ts CHANGED
@@ -29,9 +29,29 @@ export const DOC_METADATA: Record<string, DocMetadata> = {
29
  "Framework for building LLM apps: chains, agents, tool-calling, and production observability.",
30
  docsUrl: "https://docs.langchain.com/oss/python/langchain/overview",
31
  },
 
 
 
 
 
 
 
 
 
 
32
  openai_cookbooks: {
33
  description:
34
  "Example notebooks and recipes from OpenAI covering practical patterns for using their APIs.",
35
  docsUrl: "https://cookbook.openai.com",
36
  },
 
 
 
 
 
 
 
 
 
 
37
  };
 
29
  "Framework for building LLM apps: chains, agents, tool-calling, and production observability.",
30
  docsUrl: "https://docs.langchain.com/oss/python/langchain/overview",
31
  },
32
+ langgraph: {
33
+ description:
34
+ "Graph-based runtime for reliable, stateful AI agents with persistence, streaming, human review, and deployment patterns.",
35
+ docsUrl: "https://docs.langchain.com/oss/python/langgraph/overview",
36
+ },
37
+ deep_agents: {
38
+ description:
39
+ "LangChain's deep agent harness for planning, delegation, filesystem context, and longer-running agent workflows.",
40
+ docsUrl: "https://docs.langchain.com/oss/python/deepagents/overview",
41
+ },
42
  openai_cookbooks: {
43
  description:
44
  "Example notebooks and recipes from OpenAI covering practical patterns for using their APIs.",
45
  docsUrl: "https://cookbook.openai.com",
46
  },
47
+ openai_docs: {
48
+ description:
49
+ "Official OpenAI API, Agents SDK, and Codex documentation from the developer docs Markdown index.",
50
+ docsUrl: "https://developers.openai.com",
51
+ },
52
+ claude_code_docs: {
53
+ description:
54
+ "Official Claude Code and Claude Agent SDK documentation from Anthropic's Markdown index.",
55
+ docsUrl: "https://code.claude.com/docs/en/overview",
56
+ },
57
  };
scripts/chat_service.py CHANGED
@@ -19,6 +19,7 @@ from .chat_types import ChatEvent, ChatRequest, ChatTurn, SourceMatch
19
  from .chroma_rag import LocalChromaRetriever, format_tool_payload, parse_tool_payload
20
  from .prompts import build_system_prompt
21
  from .setup import (
 
22
  COURSE_SOURCE_KEYS,
23
  DOCUMENT_DICT_PATH,
24
  SOURCE_KEY_TO_LABEL,
@@ -51,6 +52,7 @@ def get_retriever() -> LocalChromaRetriever:
51
  db_path=VECTOR_DB_DIR,
52
  collection_name=VECTOR_COLLECTION_NAME,
53
  document_dict_path=DOCUMENT_DICT_PATH,
 
54
  cohere_api_key=cohere_api_key,
55
  )
56
 
 
19
  from .chroma_rag import LocalChromaRetriever, format_tool_payload, parse_tool_payload
20
  from .prompts import build_system_prompt
21
  from .setup import (
22
+ BM25_INDEX_PATH,
23
  COURSE_SOURCE_KEYS,
24
  DOCUMENT_DICT_PATH,
25
  SOURCE_KEY_TO_LABEL,
 
52
  db_path=VECTOR_DB_DIR,
53
  collection_name=VECTOR_COLLECTION_NAME,
54
  document_dict_path=DOCUMENT_DICT_PATH,
55
+ bm25_index_path=BM25_INDEX_PATH,
56
  cohere_api_key=cohere_api_key,
57
  )
58
 
scripts/chroma_rag.py CHANGED
@@ -2,7 +2,13 @@ from __future__ import annotations
2
 
3
  import json
4
  import math
 
5
  import pickle
 
 
 
 
 
6
  from dataclasses import asdict, dataclass
7
  from pathlib import Path
8
  from typing import Any, Iterable
@@ -17,13 +23,100 @@ from tqdm.auto import tqdm
17
 
18
  DEFAULT_CHUNK_SIZE = 800
19
  DEFAULT_CHUNK_OVERLAP = 100
 
20
  DEFAULT_DENSE_TOP_K = 15
 
 
21
  DEFAULT_RERANK_TOP_K = 5
 
22
  DEFAULT_CONTEXT_TOKEN_BUDGET = 100_000
23
  DEFAULT_EMBED_MODEL = "embed-v4.0"
24
  DEFAULT_RERANK_MODEL = "rerank-v4.0-fast"
25
  DEFAULT_ENCODING = "cl100k_base"
26
  DEFAULT_OUTPUT_DIMENSION = 1024
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
27
 
28
 
29
  @dataclass(slots=True)
@@ -34,6 +127,20 @@ class ChunkRecord:
34
  metadata: dict[str, Any]
35
 
36
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
37
  @dataclass(slots=True)
38
  class SearchResult:
39
  chunk_id: str
@@ -46,6 +153,108 @@ class SearchResult:
46
  score: float
47
  content: str
48
  chunk_content: str
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
49
 
50
 
51
  def batched(items: list[Any], size: int) -> Iterable[list[Any]]:
@@ -71,6 +280,133 @@ def get_token_encoding(model_name: str | None = None) -> tiktoken.Encoding:
71
  return tiktoken.get_encoding(DEFAULT_ENCODING)
72
 
73
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
74
  def token_window_chunks(
75
  text: str,
76
  chunk_size: int = DEFAULT_CHUNK_SIZE,
@@ -96,33 +432,166 @@ def token_window_chunks(
96
  return chunks
97
 
98
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
99
  def build_chunk_records(
100
  documents: list[dict[str, Any]],
101
  chunk_size: int = DEFAULT_CHUNK_SIZE,
102
  chunk_overlap: int = DEFAULT_CHUNK_OVERLAP,
103
  ) -> list[ChunkRecord]:
104
  chunk_records: list[ChunkRecord] = []
 
105
  for document in documents:
106
- chunks = token_window_chunks(
107
  document["content"],
 
108
  chunk_size=chunk_size,
109
  chunk_overlap=chunk_overlap,
110
  )
111
- for index, chunk_text in enumerate(chunks):
 
 
 
112
  metadata = {
113
  "doc_id": document["doc_id"],
114
  "title": document["name"],
115
  "url": document["url"],
116
- "source": document["source"],
 
117
  "retrieve_doc": document["retrieve_doc"],
118
  "tokens": document["tokens"],
 
119
  "chunk_index": index,
 
120
  }
121
  chunk_records.append(
122
  ChunkRecord(
123
  chunk_id=f'{document["doc_id"]}:{index}',
124
  doc_id=document["doc_id"],
125
- text=chunk_text,
126
  metadata=metadata,
127
  )
128
  )
@@ -255,6 +724,121 @@ def _title_from_url(url: str) -> str:
255
  return unquote(slug).replace("-", " ").replace("_", " ").strip()
256
 
257
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
258
  def _cohere_embeddings_list(response: Any) -> list[list[float]]:
259
  embeddings = getattr(response, "embeddings", None)
260
  if embeddings is None:
@@ -276,24 +860,96 @@ def embed_texts(
276
  input_type: str,
277
  model: str = DEFAULT_EMBED_MODEL,
278
  output_dimension: int = DEFAULT_OUTPUT_DIMENSION,
279
- batch_size: int = 96,
 
 
 
280
  show_progress: bool = False,
281
  progress_desc: str = "Embedding",
282
  ) -> list[list[float]]:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
283
  vectors: list[list[float]] = []
284
  progress = None
285
  if show_progress and texts:
286
  progress = tqdm(total=len(texts), desc=progress_desc, unit="chunk")
287
 
288
  try:
289
- for batch in batched(texts, batch_size):
290
- response = client.embed(
291
- model=model,
292
- input_type=input_type,
293
- embedding_types=["float"],
294
- output_dimension=output_dimension,
295
- texts=batch,
296
- )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
297
  vectors.extend(_cohere_embeddings_list(response))
298
  if progress is not None:
299
  progress.update(len(batch))
@@ -336,6 +992,8 @@ def rerank_results(
336
  score=float(item.relevance_score),
337
  content=result.content,
338
  chunk_content=result.chunk_content,
 
 
339
  )
340
  )
341
  return reranked
@@ -363,6 +1021,113 @@ def _flatten_query_results(values: list[list[Any]] | None) -> list[Any]:
363
  return values[0]
364
 
365
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
366
  class LocalChromaRetriever:
367
  def __init__(
368
  self,
@@ -374,7 +1139,11 @@ class LocalChromaRetriever:
374
  embed_model: str = DEFAULT_EMBED_MODEL,
375
  rerank_model: str = DEFAULT_RERANK_MODEL,
376
  dense_top_k: int = DEFAULT_DENSE_TOP_K,
 
 
377
  rerank_top_k: int = DEFAULT_RERANK_TOP_K,
 
 
378
  answer_model_name: str | None = None,
379
  token_budget: int = DEFAULT_CONTEXT_TOKEN_BUDGET,
380
  ) -> None:
@@ -382,17 +1151,24 @@ class LocalChromaRetriever:
382
  self._collection_name = collection_name
383
  self._document_dict_path = document_dict_path
384
  self._dense_top_k = dense_top_k
 
 
385
  self._rerank_top_k = rerank_top_k
 
386
  self._token_budget = token_budget
387
  self._embed_model = embed_model
388
  self._rerank_model = rerank_model
389
  self._encoding = get_token_encoding(answer_model_name)
 
 
 
390
 
391
  client = chromadb.PersistentClient(path=db_path)
392
  self._collection = client.get_or_create_collection(name=collection_name)
393
  with open(document_dict_path, "rb") as handle:
394
  self._document_dict: dict[str, dict[str, Any]] = pickle.load(handle)
395
 
 
396
  self._cohere = cohere.ClientV2(api_key=cohere_api_key)
397
 
398
  def search(
@@ -400,6 +1176,31 @@ class LocalChromaRetriever:
400
  query: str,
401
  *,
402
  allowed_sources: list[str] | None = None,
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
403
  ) -> list[SearchResult]:
404
  query_embedding = embed_texts(
405
  self._cohere,
@@ -422,78 +1223,114 @@ class LocalChromaRetriever:
422
  distances = _flatten_query_results(raw_results.get("distances"))
423
 
424
  dense_hits: list[SearchResult] = []
425
- seen_doc_ids: set[str] = set()
426
  for chunk_id, chunk_text, metadata, distance in zip(
427
  chunk_ids, documents, metadatas, distances, strict=False
428
  ):
429
  if metadata is None:
430
  continue
431
 
432
- doc_id = str(metadata["doc_id"])
433
- if doc_id in seen_doc_ids:
434
- continue
435
- seen_doc_ids.add(doc_id)
436
-
437
- full_doc = self._document_dict.get(doc_id)
438
- if metadata.get("retrieve_doc") and full_doc is not None:
439
- content = get_full_doc_content(full_doc)
440
- else:
441
- content = chunk_text
442
-
443
- full_doc_name = full_doc.get("name") if isinstance(full_doc, dict) else None
444
- url = _string_metadata_value(
445
- metadata.get("url"),
446
- full_doc.get("url") if isinstance(full_doc, dict) else None,
447
- )
448
- title = _string_metadata_value(
449
- metadata.get("title"),
450
- full_doc_name,
451
- _title_from_url(url),
452
- doc_id,
453
- )
454
- source = _string_metadata_value(
455
- metadata.get("source"),
456
- full_doc.get("source") if isinstance(full_doc, dict) else None,
457
- default="unknown",
458
- )
459
- retrieve_doc = bool(
460
- metadata.get("retrieve_doc")
461
- if "retrieve_doc" in metadata
462
- else (
463
- full_doc.get("retrieve_doc") if isinstance(full_doc, dict) else False
464
- )
465
- )
466
- tokens_value = metadata.get("tokens")
467
- if _is_missing_metadata_value(tokens_value) and isinstance(full_doc, dict):
468
- tokens_value = full_doc.get("tokens")
469
- try:
470
- tokens = int(tokens_value)
471
- except (TypeError, ValueError):
472
- tokens = 0
473
-
474
  dense_hits.append(
475
- SearchResult(
476
  chunk_id=str(chunk_id),
477
- doc_id=doc_id,
478
- title=title,
479
- url=url,
480
- source=source,
481
- retrieve_doc=retrieve_doc,
482
- tokens=tokens,
483
  score=_distance_to_score(distance),
484
- content=content,
485
- chunk_content=str(chunk_text),
 
486
  )
487
  )
 
488
 
489
- reranked = rerank_results(
490
- self._cohere,
491
- query,
492
- dense_hits,
493
- model=self._rerank_model,
494
- top_n=self._rerank_top_k,
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
495
  )
496
- return self._apply_token_budget(reranked)
497
 
498
  def _apply_token_budget(self, results: list[SearchResult]) -> list[SearchResult]:
499
  filtered: list[SearchResult] = []
 
2
 
3
  import json
4
  import math
5
+ import os
6
  import pickle
7
+ import random
8
+ import re
9
+ import threading
10
+ import time
11
+ from collections import Counter, deque
12
  from dataclasses import asdict, dataclass
13
  from pathlib import Path
14
  from typing import Any, Iterable
 
23
 
24
  DEFAULT_CHUNK_SIZE = 800
25
  DEFAULT_CHUNK_OVERLAP = 100
26
+ DEFAULT_MAX_CHUNK_TOKENS = 1200
27
  DEFAULT_DENSE_TOP_K = 15
28
+ DEFAULT_BM25_TOP_K = 30
29
+ DEFAULT_FUSION_TOP_K = 30
30
  DEFAULT_RERANK_TOP_K = 5
31
+ DEFAULT_RRF_K = 60
32
  DEFAULT_CONTEXT_TOKEN_BUDGET = 100_000
33
  DEFAULT_EMBED_MODEL = "embed-v4.0"
34
  DEFAULT_RERANK_MODEL = "rerank-v4.0-fast"
35
  DEFAULT_ENCODING = "cl100k_base"
36
  DEFAULT_OUTPUT_DIMENSION = 1024
37
+ DEFAULT_COHERE_EMBED_BATCH_SIZE = 96
38
+ DEFAULT_COHERE_EMBED_INPUTS_PER_MINUTE = 2_000
39
+ DEFAULT_COHERE_EMBED_TPM_LIMIT = 0
40
+ DEFAULT_COHERE_EMBED_RPM_LIMIT = 0
41
+ DEFAULT_COHERE_EMBED_RATE_LIMIT_MARGIN = 0.8
42
+ DEFAULT_COHERE_EMBED_WINDOW_SECONDS = 60.0
43
+ DEFAULT_COHERE_EMBED_RETRY_ATTEMPTS = 8
44
+ DEFAULT_SOURCE_VERSIONS_PATH = "data/source_versions.json"
45
+
46
+ BM25_TOKEN_RE = re.compile(r"[A-Za-z][A-Za-z0-9_./:-]*|\d+(?:\.\d+)*")
47
+ CAMEL_CASE_RE = re.compile(r"[A-Z]?[a-z]+|[A-Z]+(?=[A-Z]|$)|\d+")
48
+ MARKDOWN_HEADING_RE = re.compile(r"^(#{1,6})\s+(.+?)\s*#*\s*$")
49
+ MARKDOWN_FENCE_RE = re.compile(r"^\s*(`{3,}|~{3,})")
50
+ BM25_STOP_WORDS = frozenset(
51
+ {
52
+ "a",
53
+ "an",
54
+ "and",
55
+ "are",
56
+ "as",
57
+ "at",
58
+ "be",
59
+ "by",
60
+ "for",
61
+ "from",
62
+ "how",
63
+ "i",
64
+ "in",
65
+ "is",
66
+ "it",
67
+ "of",
68
+ "on",
69
+ "or",
70
+ "that",
71
+ "the",
72
+ "this",
73
+ "to",
74
+ "use",
75
+ "what",
76
+ "when",
77
+ "where",
78
+ "which",
79
+ "with",
80
+ "you",
81
+ "your",
82
+ }
83
+ )
84
+
85
+
86
+ class SyncWindowLimiter:
87
+ def __init__(self, units_per_window: int, window_seconds: float) -> None:
88
+ self.units_per_window = max(1, units_per_window)
89
+ self.window_seconds = window_seconds
90
+ self._events: deque[tuple[float, int]] = deque()
91
+ self._used_units = 0
92
+ self._lock = threading.Lock()
93
+
94
+ def acquire(self, units: int) -> None:
95
+ units = min(max(1, units), self.units_per_window)
96
+
97
+ while True:
98
+ with self._lock:
99
+ now = time.monotonic()
100
+ self._prune(now)
101
+
102
+ if self._used_units + units <= self.units_per_window:
103
+ self._events.append((now, units))
104
+ self._used_units += units
105
+ return
106
+
107
+ oldest_at, _ = self._events[0]
108
+ delay = max(0.1, self.window_seconds - (now - oldest_at))
109
+
110
+ time.sleep(delay + random.uniform(0.1, 0.75))
111
+
112
+ def _prune(self, now: float) -> None:
113
+ while self._events and now - self._events[0][0] >= self.window_seconds:
114
+ _, units = self._events.popleft()
115
+ self._used_units -= units
116
+
117
+
118
+ _cohere_limiter_lock = threading.Lock()
119
+ _cohere_limiters: dict[tuple[str, int, float], SyncWindowLimiter] = {}
120
 
121
 
122
  @dataclass(slots=True)
 
127
  metadata: dict[str, Any]
128
 
129
 
130
+ @dataclass(slots=True)
131
+ class MarkdownUnit:
132
+ text: str
133
+ heading_path: tuple[str, ...]
134
+ is_code: bool = False
135
+
136
+
137
+ @dataclass(slots=True)
138
+ class MarkdownChunk:
139
+ text: str
140
+ heading_path: tuple[str, ...]
141
+ tokens: int
142
+
143
+
144
  @dataclass(slots=True)
145
  class SearchResult:
146
  chunk_id: str
 
153
  score: float
154
  content: str
155
  chunk_content: str
156
+ heading_path: str = ""
157
+ retrieval_method: str = ""
158
+
159
+
160
+ @dataclass(slots=True)
161
+ class BM25Index:
162
+ records: list[ChunkRecord]
163
+ postings: dict[str, list[tuple[int, int]]]
164
+ document_frequencies: dict[str, int]
165
+ document_lengths: list[int]
166
+ average_document_length: float
167
+ k1: float = 1.5
168
+ b: float = 0.75
169
+
170
+ @classmethod
171
+ def build(
172
+ cls,
173
+ records: list[ChunkRecord],
174
+ *,
175
+ k1: float = 1.5,
176
+ b: float = 0.75,
177
+ ) -> "BM25Index":
178
+ postings: dict[str, list[tuple[int, int]]] = {}
179
+ document_frequencies: dict[str, int] = {}
180
+ document_lengths: list[int] = []
181
+
182
+ for doc_index, record in enumerate(records):
183
+ terms = tokenize_for_bm25(
184
+ format_chunk_for_retrieval(record.text, record.metadata)
185
+ )
186
+ counts = Counter(terms)
187
+ document_lengths.append(sum(counts.values()))
188
+
189
+ for term, term_frequency in counts.items():
190
+ postings.setdefault(term, []).append((doc_index, term_frequency))
191
+ for term in counts:
192
+ document_frequencies[term] = document_frequencies.get(term, 0) + 1
193
+
194
+ average_document_length = (
195
+ sum(document_lengths) / len(document_lengths) if document_lengths else 0.0
196
+ )
197
+ return cls(
198
+ records=records,
199
+ postings=postings,
200
+ document_frequencies=document_frequencies,
201
+ document_lengths=document_lengths,
202
+ average_document_length=average_document_length,
203
+ k1=k1,
204
+ b=b,
205
+ )
206
+
207
+ def search(
208
+ self,
209
+ query: str,
210
+ *,
211
+ allowed_sources: list[str] | None = None,
212
+ top_k: int = DEFAULT_BM25_TOP_K,
213
+ ) -> list[tuple[ChunkRecord, float]]:
214
+ query_terms = tokenize_for_bm25(query)
215
+ if not query_terms or not self.records:
216
+ return []
217
+
218
+ allowed = set(allowed_sources or [])
219
+ total_documents = len(self.records)
220
+ scores: dict[int, float] = {}
221
+
222
+ for term in set(query_terms):
223
+ postings = self.postings.get(term)
224
+ if not postings:
225
+ continue
226
+
227
+ document_frequency = self.document_frequencies.get(term, len(postings))
228
+ idf = math.log(
229
+ 1.0
230
+ + (total_documents - document_frequency + 0.5)
231
+ / (document_frequency + 0.5)
232
+ )
233
+
234
+ for doc_index, term_frequency in postings:
235
+ record = self.records[doc_index]
236
+ if allowed and record.metadata.get("source") not in allowed:
237
+ continue
238
+
239
+ document_length = self.document_lengths[doc_index]
240
+ if document_length <= 0 or self.average_document_length <= 0:
241
+ continue
242
+
243
+ denominator = term_frequency + self.k1 * (
244
+ 1.0
245
+ - self.b
246
+ + self.b * document_length / self.average_document_length
247
+ )
248
+ scores[doc_index] = scores.get(doc_index, 0.0) + idf * (
249
+ term_frequency * (self.k1 + 1.0) / denominator
250
+ )
251
+
252
+ ranked = sorted(scores.items(), key=lambda item: item[1], reverse=True)
253
+ return [
254
+ (self.records[doc_index], score)
255
+ for doc_index, score in ranked[: max(0, top_k)]
256
+ if score > 0.0
257
+ ]
258
 
259
 
260
  def batched(items: list[Any], size: int) -> Iterable[list[Any]]:
 
280
  return tiktoken.get_encoding(DEFAULT_ENCODING)
281
 
282
 
283
+ def load_source_versions(
284
+ path: str = DEFAULT_SOURCE_VERSIONS_PATH,
285
+ ) -> dict[str, dict[str, Any]]:
286
+ if not os.path.exists(path):
287
+ return {}
288
+ try:
289
+ with open(path, "r", encoding="utf-8") as handle:
290
+ data = json.load(handle)
291
+ except (OSError, json.JSONDecodeError):
292
+ return {}
293
+ if not isinstance(data, dict):
294
+ return {}
295
+ return {
296
+ str(source): dict(metadata)
297
+ for source, metadata in data.items()
298
+ if isinstance(metadata, dict)
299
+ }
300
+
301
+
302
+ def source_version_for(
303
+ source: str,
304
+ source_versions: dict[str, dict[str, Any]] | None = None,
305
+ ) -> str:
306
+ metadata = (source_versions or {}).get(source, {})
307
+ for key in ("version", "sha", "indexedAt"):
308
+ value = metadata.get(key)
309
+ if value:
310
+ return str(value)
311
+ return ""
312
+
313
+
314
+ def clean_heading_text(text: str) -> str:
315
+ return re.sub(r"\s+", " ", text.strip().strip("#").strip())
316
+
317
+
318
+ def parse_markdown_units(
319
+ text: str,
320
+ *,
321
+ default_heading_path: tuple[str, ...] = (),
322
+ ) -> list[MarkdownUnit]:
323
+ lines = text.splitlines()
324
+ units: list[MarkdownUnit] = []
325
+ heading_stack: list[str] = []
326
+ text_buffer: list[str] = []
327
+
328
+ def active_heading_path() -> tuple[str, ...]:
329
+ return tuple(heading_stack) or default_heading_path
330
+
331
+ def flush_text_buffer() -> None:
332
+ nonlocal text_buffer
333
+ paragraph: list[str] = []
334
+ for buffered_line in text_buffer:
335
+ if buffered_line.strip():
336
+ paragraph.append(buffered_line)
337
+ continue
338
+ if paragraph:
339
+ units.append(
340
+ MarkdownUnit(
341
+ text="\n".join(paragraph).strip(),
342
+ heading_path=active_heading_path(),
343
+ )
344
+ )
345
+ paragraph = []
346
+ if paragraph:
347
+ units.append(
348
+ MarkdownUnit(
349
+ text="\n".join(paragraph).strip(),
350
+ heading_path=active_heading_path(),
351
+ )
352
+ )
353
+ text_buffer = []
354
+
355
+ index = 0
356
+ while index < len(lines):
357
+ line = lines[index]
358
+ fence_match = MARKDOWN_FENCE_RE.match(line)
359
+ if fence_match:
360
+ flush_text_buffer()
361
+ fence = fence_match.group(1)
362
+ fence_char = fence[0]
363
+ fence_len = len(fence)
364
+ code_lines = [line]
365
+ index += 1
366
+ while index < len(lines):
367
+ code_line = lines[index]
368
+ code_lines.append(code_line)
369
+ close_match = MARKDOWN_FENCE_RE.match(code_line)
370
+ if (
371
+ close_match
372
+ and close_match.group(1)[0] == fence_char
373
+ and len(close_match.group(1)) >= fence_len
374
+ ):
375
+ index += 1
376
+ break
377
+ index += 1
378
+ units.append(
379
+ MarkdownUnit(
380
+ text="\n".join(code_lines).strip(),
381
+ heading_path=active_heading_path(),
382
+ is_code=True,
383
+ )
384
+ )
385
+ continue
386
+
387
+ heading_match = MARKDOWN_HEADING_RE.match(line)
388
+ if heading_match:
389
+ flush_text_buffer()
390
+ level = len(heading_match.group(1))
391
+ title = clean_heading_text(heading_match.group(2))
392
+ heading_stack = heading_stack[: level - 1]
393
+ heading_stack.append(title)
394
+ units.append(
395
+ MarkdownUnit(
396
+ text=line.strip(),
397
+ heading_path=active_heading_path(),
398
+ )
399
+ )
400
+ index += 1
401
+ continue
402
+
403
+ text_buffer.append(line)
404
+ index += 1
405
+
406
+ flush_text_buffer()
407
+ return [unit for unit in units if unit.text.strip()]
408
+
409
+
410
  def token_window_chunks(
411
  text: str,
412
  chunk_size: int = DEFAULT_CHUNK_SIZE,
 
432
  return chunks
433
 
434
 
435
+ def _heading_path_text(heading_path: tuple[str, ...] | list[str] | str | None) -> str:
436
+ if not heading_path:
437
+ return ""
438
+ if isinstance(heading_path, str):
439
+ return heading_path
440
+ return " > ".join(str(part) for part in heading_path if str(part).strip())
441
+
442
+
443
+ def _split_large_unit(
444
+ unit: MarkdownUnit,
445
+ *,
446
+ encoding: tiktoken.Encoding,
447
+ chunk_size: int,
448
+ chunk_overlap: int,
449
+ ) -> list[MarkdownUnit]:
450
+ if unit.is_code:
451
+ return [unit]
452
+
453
+ token_count = len(encoding.encode(unit.text, disallowed_special=()))
454
+ if token_count <= chunk_size:
455
+ return [unit]
456
+ return [
457
+ MarkdownUnit(text=chunk, heading_path=unit.heading_path)
458
+ for chunk in token_window_chunks(
459
+ unit.text,
460
+ chunk_size=chunk_size,
461
+ chunk_overlap=chunk_overlap,
462
+ )
463
+ ]
464
+
465
+
466
+ def heading_aware_markdown_chunks(
467
+ text: str,
468
+ *,
469
+ title: str = "",
470
+ chunk_size: int = DEFAULT_CHUNK_SIZE,
471
+ chunk_overlap: int = DEFAULT_CHUNK_OVERLAP,
472
+ encoding_name: str = DEFAULT_ENCODING,
473
+ ) -> list[MarkdownChunk]:
474
+ encoding = tiktoken.get_encoding(encoding_name)
475
+ default_heading_path = (title,) if title else ()
476
+ units = parse_markdown_units(text, default_heading_path=default_heading_path)
477
+ if not units:
478
+ return []
479
+
480
+ budget_size = min(max(chunk_size, 1), DEFAULT_MAX_CHUNK_TOKENS)
481
+ overlap_size = chunk_overlap
482
+
483
+ def unit_size(value: str) -> int:
484
+ return len(encoding.encode(value, disallowed_special=()))
485
+
486
+ chunks: list[MarkdownChunk] = []
487
+ current_parts: list[str] = []
488
+ current_heading_path: tuple[str, ...] = ()
489
+ current_size = 0
490
+
491
+ def flush_current() -> None:
492
+ nonlocal current_parts, current_heading_path, current_size
493
+ text_value = "\n\n".join(part for part in current_parts if part.strip()).strip()
494
+ if text_value:
495
+ chunks.append(
496
+ MarkdownChunk(
497
+ text=text_value,
498
+ heading_path=current_heading_path,
499
+ tokens=len(encoding.encode(text_value, disallowed_special=())),
500
+ )
501
+ )
502
+ current_parts = []
503
+ current_heading_path = ()
504
+ current_size = 0
505
+
506
+ for unit in units:
507
+ split_units = _split_large_unit(
508
+ unit,
509
+ encoding=encoding,
510
+ chunk_size=budget_size,
511
+ chunk_overlap=overlap_size,
512
+ )
513
+ for split_unit in split_units:
514
+ size = unit_size(split_unit.text)
515
+ heading_changed = (
516
+ current_heading_path
517
+ and split_unit.heading_path != current_heading_path
518
+ )
519
+ would_exceed = current_parts and current_size + size > budget_size
520
+ if heading_changed or would_exceed:
521
+ flush_current()
522
+
523
+ if not current_parts:
524
+ current_heading_path = split_unit.heading_path
525
+ current_parts.append(split_unit.text)
526
+ current_size += size
527
+
528
+ flush_current()
529
+ return chunks
530
+
531
+
532
+ def build_chunk_retrieval_header(metadata: dict[str, Any]) -> str:
533
+ lines: list[str] = []
534
+ title = _string_metadata_value(metadata.get("title"), metadata.get("name"))
535
+ source = _string_metadata_value(metadata.get("source"))
536
+ version = _string_metadata_value(metadata.get("source_version"))
537
+ heading_path = _string_metadata_value(metadata.get("heading_path"))
538
+
539
+ if title:
540
+ lines.append(f"Title: {title}")
541
+ if source:
542
+ lines.append(f"Source: {source}")
543
+ if version:
544
+ lines.append(f"Version: {version}")
545
+ if heading_path:
546
+ lines.append(f"Heading path: {heading_path}")
547
+ return "\n".join(lines)
548
+
549
+
550
+ def format_chunk_for_retrieval(text: str, metadata: dict[str, Any]) -> str:
551
+ text = text.strip()
552
+ header = build_chunk_retrieval_header(metadata)
553
+ if not header:
554
+ return text
555
+ if not text:
556
+ return header
557
+ return f"{header}\n\n{text}"
558
+
559
+
560
  def build_chunk_records(
561
  documents: list[dict[str, Any]],
562
  chunk_size: int = DEFAULT_CHUNK_SIZE,
563
  chunk_overlap: int = DEFAULT_CHUNK_OVERLAP,
564
  ) -> list[ChunkRecord]:
565
  chunk_records: list[ChunkRecord] = []
566
+ source_versions = load_source_versions()
567
  for document in documents:
568
+ chunks = heading_aware_markdown_chunks(
569
  document["content"],
570
+ title=str(document.get("name") or ""),
571
  chunk_size=chunk_size,
572
  chunk_overlap=chunk_overlap,
573
  )
574
+ source = str(document["source"])
575
+ source_version = source_version_for(source, source_versions)
576
+ for index, chunk in enumerate(chunks):
577
+ heading_path = _heading_path_text(chunk.heading_path)
578
  metadata = {
579
  "doc_id": document["doc_id"],
580
  "title": document["name"],
581
  "url": document["url"],
582
+ "source": source,
583
+ "source_version": source_version,
584
  "retrieve_doc": document["retrieve_doc"],
585
  "tokens": document["tokens"],
586
+ "chunk_tokens": chunk.tokens,
587
  "chunk_index": index,
588
+ "heading_path": heading_path,
589
  }
590
  chunk_records.append(
591
  ChunkRecord(
592
  chunk_id=f'{document["doc_id"]}:{index}',
593
  doc_id=document["doc_id"],
594
+ text=chunk.text,
595
  metadata=metadata,
596
  )
597
  )
 
724
  return unquote(slug).replace("-", " ").replace("_", " ").strip()
725
 
726
 
727
+ def _env_int(name: str, default: int) -> int:
728
+ value = os.getenv(name)
729
+ if value is None:
730
+ return default
731
+ try:
732
+ return int(value)
733
+ except ValueError:
734
+ return default
735
+
736
+
737
+ def _env_float(name: str, default: float) -> float:
738
+ value = os.getenv(name)
739
+ if value is None:
740
+ return default
741
+ try:
742
+ return float(value)
743
+ except ValueError:
744
+ return default
745
+
746
+
747
+ def _cohere_rate_limited_units(
748
+ limit: int | None,
749
+ env_name: str,
750
+ default: int,
751
+ margin: float,
752
+ ) -> int | None:
753
+ configured_limit = _env_int(env_name, default) if limit is None else limit
754
+ if configured_limit <= 0:
755
+ return None
756
+ return max(1, int(configured_limit * margin))
757
+
758
+
759
+ def _get_cohere_limiter(
760
+ name: str,
761
+ units_per_window: int,
762
+ window_seconds: float,
763
+ ) -> SyncWindowLimiter:
764
+ key = (name, units_per_window, window_seconds)
765
+ with _cohere_limiter_lock:
766
+ limiter = _cohere_limiters.get(key)
767
+ if limiter is None:
768
+ limiter = SyncWindowLimiter(units_per_window, window_seconds)
769
+ _cohere_limiters[key] = limiter
770
+ return limiter
771
+
772
+
773
+ def _count_embed_tokens(text: str, encoding: tiktoken.Encoding) -> int:
774
+ return len(encoding.encode(text, disallowed_special=()))
775
+
776
+
777
+ def _iter_cohere_embed_batches(
778
+ texts: list[str],
779
+ token_counts: list[int],
780
+ *,
781
+ batch_size: int,
782
+ max_batch_tokens: int | None,
783
+ ) -> Iterable[tuple[list[str], int]]:
784
+ batch: list[str] = []
785
+ batch_tokens = 0
786
+
787
+ for text, token_count in zip(texts, token_counts, strict=True):
788
+ should_flush = len(batch) >= batch_size
789
+ if max_batch_tokens is not None and batch:
790
+ should_flush = should_flush or batch_tokens + token_count > max_batch_tokens
791
+
792
+ if should_flush:
793
+ yield batch, batch_tokens
794
+ batch = []
795
+ batch_tokens = 0
796
+
797
+ batch.append(text)
798
+ batch_tokens += max(1, token_count)
799
+
800
+ if batch:
801
+ yield batch, batch_tokens
802
+
803
+
804
+ def _is_cohere_rate_limit_error(exc: BaseException) -> bool:
805
+ status_code = getattr(exc, "status_code", None)
806
+ if status_code == 429:
807
+ return True
808
+ if exc.__class__.__name__ == "TooManyRequestsError":
809
+ return True
810
+ return "rate limit" in str(exc).lower() and "429" in str(exc)
811
+
812
+
813
+ def _cohere_retry_after_seconds(exc: BaseException) -> float | None:
814
+ headers = getattr(exc, "headers", None)
815
+ if not isinstance(headers, dict):
816
+ return None
817
+
818
+ retry_after = headers.get("retry-after") or headers.get("Retry-After")
819
+ if retry_after is None:
820
+ return None
821
+
822
+ try:
823
+ return float(retry_after)
824
+ except ValueError:
825
+ return None
826
+
827
+
828
+ def _wait_for_cohere_retry(
829
+ exc: BaseException,
830
+ attempt: int,
831
+ window_seconds: float,
832
+ ) -> None:
833
+ retry_after = _cohere_retry_after_seconds(exc)
834
+ if retry_after is not None:
835
+ delay = retry_after
836
+ else:
837
+ delay = min(window_seconds, max(15.0, 2.0 ** attempt))
838
+
839
+ time.sleep(delay + random.uniform(0.5, 2.0))
840
+
841
+
842
  def _cohere_embeddings_list(response: Any) -> list[list[float]]:
843
  embeddings = getattr(response, "embeddings", None)
844
  if embeddings is None:
 
860
  input_type: str,
861
  model: str = DEFAULT_EMBED_MODEL,
862
  output_dimension: int = DEFAULT_OUTPUT_DIMENSION,
863
+ batch_size: int = DEFAULT_COHERE_EMBED_BATCH_SIZE,
864
+ max_inputs_per_minute: int | None = None,
865
+ max_tokens_per_minute: int | None = None,
866
+ max_requests_per_minute: int | None = None,
867
  show_progress: bool = False,
868
  progress_desc: str = "Embedding",
869
  ) -> list[list[float]]:
870
+ batch_size = max(1, batch_size)
871
+ rate_limit_margin = _env_float(
872
+ "COHERE_EMBED_RATE_LIMIT_MARGIN",
873
+ DEFAULT_COHERE_EMBED_RATE_LIMIT_MARGIN,
874
+ )
875
+ window_seconds = _env_float(
876
+ "COHERE_EMBED_WINDOW_SECONDS",
877
+ DEFAULT_COHERE_EMBED_WINDOW_SECONDS,
878
+ )
879
+ token_window = _cohere_rate_limited_units(
880
+ max_tokens_per_minute,
881
+ "COHERE_EMBED_TPM_LIMIT",
882
+ DEFAULT_COHERE_EMBED_TPM_LIMIT,
883
+ rate_limit_margin,
884
+ )
885
+ input_window = _cohere_rate_limited_units(
886
+ max_inputs_per_minute,
887
+ "COHERE_EMBED_INPUTS_PER_MINUTE",
888
+ DEFAULT_COHERE_EMBED_INPUTS_PER_MINUTE,
889
+ rate_limit_margin,
890
+ )
891
+ request_window = _cohere_rate_limited_units(
892
+ max_requests_per_minute,
893
+ "COHERE_EMBED_RPM_LIMIT",
894
+ DEFAULT_COHERE_EMBED_RPM_LIMIT,
895
+ rate_limit_margin,
896
+ )
897
+ retry_attempts = max(
898
+ 1,
899
+ _env_int("COHERE_EMBED_RETRY_ATTEMPTS", DEFAULT_COHERE_EMBED_RETRY_ATTEMPTS),
900
+ )
901
+ encoding = tiktoken.get_encoding(DEFAULT_ENCODING)
902
+ token_counts = [_count_embed_tokens(text, encoding) for text in texts]
903
+ token_limiter = (
904
+ _get_cohere_limiter("tokens", token_window, window_seconds)
905
+ if token_window is not None
906
+ else None
907
+ )
908
+ input_limiter = (
909
+ _get_cohere_limiter("inputs", input_window, window_seconds)
910
+ if input_window is not None
911
+ else None
912
+ )
913
+ request_limiter = (
914
+ _get_cohere_limiter("requests", request_window, window_seconds)
915
+ if request_window is not None
916
+ else None
917
+ )
918
+
919
  vectors: list[list[float]] = []
920
  progress = None
921
  if show_progress and texts:
922
  progress = tqdm(total=len(texts), desc=progress_desc, unit="chunk")
923
 
924
  try:
925
+ for batch, batch_tokens in _iter_cohere_embed_batches(
926
+ texts,
927
+ token_counts,
928
+ batch_size=batch_size,
929
+ max_batch_tokens=token_window,
930
+ ):
931
+ for attempt in range(1, retry_attempts + 1):
932
+ if token_limiter is not None:
933
+ token_limiter.acquire(batch_tokens)
934
+ if input_limiter is not None:
935
+ input_limiter.acquire(len(batch))
936
+ if request_limiter is not None:
937
+ request_limiter.acquire(1)
938
+
939
+ try:
940
+ response = client.embed(
941
+ model=model,
942
+ input_type=input_type,
943
+ embedding_types=["float"],
944
+ output_dimension=output_dimension,
945
+ texts=batch,
946
+ )
947
+ break
948
+ except Exception as exc:
949
+ if not _is_cohere_rate_limit_error(exc) or attempt == retry_attempts:
950
+ raise
951
+ _wait_for_cohere_retry(exc, attempt, window_seconds)
952
+
953
  vectors.extend(_cohere_embeddings_list(response))
954
  if progress is not None:
955
  progress.update(len(batch))
 
992
  score=float(item.relevance_score),
993
  content=result.content,
994
  chunk_content=result.chunk_content,
995
+ heading_path=result.heading_path,
996
+ retrieval_method=result.retrieval_method,
997
  )
998
  )
999
  return reranked
 
1021
  return values[0]
1022
 
1023
 
1024
+ def _add_bm25_token(tokens: list[str], token: str) -> None:
1025
+ token = token.lower().strip("._-/:-")
1026
+ if not token:
1027
+ return
1028
+ if token in BM25_STOP_WORDS:
1029
+ return
1030
+ if len(token) == 1 and token not in {"c", "r"}:
1031
+ return
1032
+ tokens.append(token)
1033
+
1034
+
1035
+ def tokenize_for_bm25(text: str) -> list[str]:
1036
+ tokens: list[str] = []
1037
+ for match in BM25_TOKEN_RE.finditer(text):
1038
+ raw_token = match.group(0)
1039
+ _add_bm25_token(tokens, raw_token)
1040
+
1041
+ for part in re.split(r"[._/\-:]+", raw_token):
1042
+ _add_bm25_token(tokens, part)
1043
+ for camel_part in CAMEL_CASE_RE.findall(part):
1044
+ _add_bm25_token(tokens, camel_part)
1045
+
1046
+ return tokens
1047
+
1048
+
1049
+ def save_bm25_index(index: BM25Index, output_file: str) -> None:
1050
+ ensure_parent_dir(output_file)
1051
+ with open(output_file, "wb") as handle:
1052
+ pickle.dump(index, handle)
1053
+
1054
+
1055
+ def load_bm25_index(path: str) -> BM25Index | None:
1056
+ if not os.path.exists(path):
1057
+ return None
1058
+ with open(path, "rb") as handle:
1059
+ index = pickle.load(handle)
1060
+ if isinstance(index, BM25Index):
1061
+ return index
1062
+ return None
1063
+
1064
+
1065
+ def default_bm25_index_path(document_dict_path: str) -> str:
1066
+ path = Path(document_dict_path)
1067
+ name = path.name
1068
+ if name.startswith("document_dict_") and name.endswith(".pkl"):
1069
+ source = name.removeprefix("document_dict_").removesuffix(".pkl")
1070
+ return str(path.with_name(f"bm25_index_{source}.pkl"))
1071
+ return str(path.with_name("bm25_index.pkl"))
1072
+
1073
+
1074
+ def result_dedupe_key(result: SearchResult) -> str:
1075
+ if result.heading_path:
1076
+ return f"{result.doc_id}::{result.heading_path}"
1077
+ return result.doc_id or result.chunk_id
1078
+
1079
+
1080
+ def reciprocal_rank_fusion(
1081
+ ranked_lists: list[list[SearchResult]],
1082
+ *,
1083
+ rrf_k: int = DEFAULT_RRF_K,
1084
+ top_k: int = DEFAULT_FUSION_TOP_K,
1085
+ ) -> list[SearchResult]:
1086
+ fused_scores: dict[str, float] = {}
1087
+ representatives: dict[str, SearchResult] = {}
1088
+
1089
+ for ranked_results in ranked_lists:
1090
+ for rank, result in enumerate(ranked_results, start=1):
1091
+ key = result_dedupe_key(result)
1092
+ fused_scores[key] = fused_scores.get(key, 0.0) + 1.0 / (rrf_k + rank)
1093
+
1094
+ current = representatives.get(key)
1095
+ if current is None:
1096
+ representatives[key] = result
1097
+ elif result.retrieval_method == "bm25" and current.retrieval_method != "bm25":
1098
+ representatives[key] = result
1099
+ elif (
1100
+ result.retrieval_method == current.retrieval_method
1101
+ and result.score > current.score
1102
+ ):
1103
+ representatives[key] = result
1104
+
1105
+ fused_results: list[SearchResult] = []
1106
+ for key, score in fused_scores.items():
1107
+ result = representatives[key]
1108
+ fused_results.append(
1109
+ SearchResult(
1110
+ chunk_id=result.chunk_id,
1111
+ doc_id=result.doc_id,
1112
+ title=result.title,
1113
+ url=result.url,
1114
+ source=result.source,
1115
+ retrieve_doc=result.retrieve_doc,
1116
+ tokens=result.tokens,
1117
+ score=score,
1118
+ content=result.content,
1119
+ chunk_content=result.chunk_content,
1120
+ heading_path=result.heading_path,
1121
+ retrieval_method="hybrid"
1122
+ if len(ranked_lists) > 1
1123
+ else result.retrieval_method,
1124
+ )
1125
+ )
1126
+
1127
+ fused_results.sort(key=lambda result: result.score, reverse=True)
1128
+ return fused_results[: max(0, top_k)]
1129
+
1130
+
1131
  class LocalChromaRetriever:
1132
  def __init__(
1133
  self,
 
1139
  embed_model: str = DEFAULT_EMBED_MODEL,
1140
  rerank_model: str = DEFAULT_RERANK_MODEL,
1141
  dense_top_k: int = DEFAULT_DENSE_TOP_K,
1142
+ bm25_top_k: int = DEFAULT_BM25_TOP_K,
1143
+ fusion_top_k: int = DEFAULT_FUSION_TOP_K,
1144
  rerank_top_k: int = DEFAULT_RERANK_TOP_K,
1145
+ rrf_k: int = DEFAULT_RRF_K,
1146
+ bm25_index_path: str | None = None,
1147
  answer_model_name: str | None = None,
1148
  token_budget: int = DEFAULT_CONTEXT_TOKEN_BUDGET,
1149
  ) -> None:
 
1151
  self._collection_name = collection_name
1152
  self._document_dict_path = document_dict_path
1153
  self._dense_top_k = dense_top_k
1154
+ self._bm25_top_k = bm25_top_k
1155
+ self._fusion_top_k = fusion_top_k
1156
  self._rerank_top_k = rerank_top_k
1157
+ self._rrf_k = rrf_k
1158
  self._token_budget = token_budget
1159
  self._embed_model = embed_model
1160
  self._rerank_model = rerank_model
1161
  self._encoding = get_token_encoding(answer_model_name)
1162
+ self._bm25_index_path = bm25_index_path or default_bm25_index_path(
1163
+ document_dict_path
1164
+ )
1165
 
1166
  client = chromadb.PersistentClient(path=db_path)
1167
  self._collection = client.get_or_create_collection(name=collection_name)
1168
  with open(document_dict_path, "rb") as handle:
1169
  self._document_dict: dict[str, dict[str, Any]] = pickle.load(handle)
1170
 
1171
+ self._bm25_index = load_bm25_index(self._bm25_index_path)
1172
  self._cohere = cohere.ClientV2(api_key=cohere_api_key)
1173
 
1174
  def search(
 
1176
  query: str,
1177
  *,
1178
  allowed_sources: list[str] | None = None,
1179
+ ) -> list[SearchResult]:
1180
+ dense_hits = self._dense_search(query, allowed_sources=allowed_sources)
1181
+ bm25_hits = self._bm25_search(query, allowed_sources=allowed_sources)
1182
+ fused_hits = reciprocal_rank_fusion(
1183
+ [hits for hits in (dense_hits, bm25_hits) if hits],
1184
+ rrf_k=self._rrf_k,
1185
+ top_k=self._fusion_top_k,
1186
+ )
1187
+ if not fused_hits:
1188
+ return []
1189
+
1190
+ reranked = rerank_results(
1191
+ self._cohere,
1192
+ query,
1193
+ fused_hits,
1194
+ model=self._rerank_model,
1195
+ top_n=self._rerank_top_k,
1196
+ )
1197
+ return self._apply_token_budget(reranked)
1198
+
1199
+ def _dense_search(
1200
+ self,
1201
+ query: str,
1202
+ *,
1203
+ allowed_sources: list[str] | None = None,
1204
  ) -> list[SearchResult]:
1205
  query_embedding = embed_texts(
1206
  self._cohere,
 
1223
  distances = _flatten_query_results(raw_results.get("distances"))
1224
 
1225
  dense_hits: list[SearchResult] = []
 
1226
  for chunk_id, chunk_text, metadata, distance in zip(
1227
  chunk_ids, documents, metadatas, distances, strict=False
1228
  ):
1229
  if metadata is None:
1230
  continue
1231
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1232
  dense_hits.append(
1233
+ self._search_result_from_metadata(
1234
  chunk_id=str(chunk_id),
 
 
 
 
 
 
1235
  score=_distance_to_score(distance),
1236
+ chunk_text=str(chunk_text),
1237
+ metadata=dict(metadata),
1238
+ retrieval_method="dense",
1239
  )
1240
  )
1241
+ return dense_hits
1242
 
1243
+ def _bm25_search(
1244
+ self,
1245
+ query: str,
1246
+ *,
1247
+ allowed_sources: list[str] | None = None,
1248
+ ) -> list[SearchResult]:
1249
+ if self._bm25_index is None:
1250
+ return []
1251
+
1252
+ return [
1253
+ self._search_result_from_metadata(
1254
+ chunk_id=record.chunk_id,
1255
+ score=score,
1256
+ chunk_text=format_chunk_for_retrieval(record.text, record.metadata),
1257
+ metadata=record.metadata,
1258
+ raw_chunk_text=record.text,
1259
+ retrieval_method="bm25",
1260
+ )
1261
+ for record, score in self._bm25_index.search(
1262
+ query,
1263
+ allowed_sources=allowed_sources,
1264
+ top_k=self._bm25_top_k,
1265
+ )
1266
+ ]
1267
+
1268
+ def _search_result_from_metadata(
1269
+ self,
1270
+ *,
1271
+ chunk_id: str,
1272
+ score: float,
1273
+ chunk_text: str,
1274
+ metadata: dict[str, Any],
1275
+ retrieval_method: str,
1276
+ raw_chunk_text: str | None = None,
1277
+ ) -> SearchResult:
1278
+ doc_id = str(metadata["doc_id"])
1279
+ full_doc = self._document_dict.get(doc_id)
1280
+ full_doc_name = full_doc.get("name") if isinstance(full_doc, dict) else None
1281
+ url = _string_metadata_value(
1282
+ metadata.get("url"),
1283
+ full_doc.get("url") if isinstance(full_doc, dict) else None,
1284
+ )
1285
+ title = _string_metadata_value(
1286
+ metadata.get("title"),
1287
+ full_doc_name,
1288
+ _title_from_url(url),
1289
+ doc_id,
1290
+ )
1291
+ source = _string_metadata_value(
1292
+ metadata.get("source"),
1293
+ full_doc.get("source") if isinstance(full_doc, dict) else None,
1294
+ default="unknown",
1295
+ )
1296
+ retrieve_doc = bool(
1297
+ metadata.get("retrieve_doc")
1298
+ if "retrieve_doc" in metadata
1299
+ else (full_doc.get("retrieve_doc") if isinstance(full_doc, dict) else False)
1300
+ )
1301
+ tokens_value = metadata.get("tokens")
1302
+ if _is_missing_metadata_value(tokens_value) and isinstance(full_doc, dict):
1303
+ tokens_value = full_doc.get("tokens")
1304
+ try:
1305
+ tokens = int(tokens_value)
1306
+ except (TypeError, ValueError):
1307
+ tokens = 0
1308
+
1309
+ exact_chunk = _string_metadata_value(
1310
+ raw_chunk_text,
1311
+ metadata.get("raw_text"),
1312
+ default=str(chunk_text),
1313
+ )
1314
+ chunk_for_context = format_chunk_for_retrieval(exact_chunk, metadata)
1315
+ if retrieve_doc and full_doc is not None:
1316
+ content = get_full_doc_content(full_doc)
1317
+ else:
1318
+ content = chunk_for_context
1319
+
1320
+ return SearchResult(
1321
+ chunk_id=chunk_id,
1322
+ doc_id=doc_id,
1323
+ title=title,
1324
+ url=url,
1325
+ source=source,
1326
+ retrieve_doc=retrieve_doc,
1327
+ tokens=tokens,
1328
+ score=score,
1329
+ content=content,
1330
+ chunk_content=exact_chunk,
1331
+ heading_path=_string_metadata_value(metadata.get("heading_path")),
1332
+ retrieval_method=retrieval_method,
1333
  )
 
1334
 
1335
  def _apply_token_budget(self, results: list[SearchResult]) -> list[SearchResult]:
1336
  filtered: list[SearchResult] = []
scripts/setup.py CHANGED
@@ -3,6 +3,15 @@ import os
3
  import logfire
4
  from dotenv import load_dotenv
5
 
 
 
 
 
 
 
 
 
 
6
  from .utils import init_mongo_db
7
 
8
  load_dotenv(override=True)
@@ -14,6 +23,7 @@ except Exception:
14
  VECTOR_DB_DIR = "data/chroma-db-all_sources"
15
  VECTOR_COLLECTION_NAME = "chroma-db-all_sources"
16
  DOCUMENT_DICT_PATH = f"{VECTOR_DB_DIR}/document_dict_all_sources.pkl"
 
17
  DEFAULT_MODEL_NAME = "google-genai:gemini-3-flash-preview"
18
 
19
  AVAILABLE_MODELS: tuple[dict[str, str], ...] = (
@@ -21,72 +31,6 @@ AVAILABLE_MODELS: tuple[dict[str, str], ...] = (
21
  {"id": "anthropic:claude-haiku-4-5", "label": "Claude Haiku 4.5"},
22
  )
23
 
24
- AVAILABLE_SOURCES_UI = [
25
- "LangChain Docs",
26
- "LlamaIndex Docs",
27
- "Transformers Docs",
28
- "PEFT Docs",
29
- "TRL Docs",
30
- "OpenAI Cookbooks",
31
- "Agentic AI Engineering",
32
- "Master AI For Work",
33
- "Full Stack AI Engineering",
34
- "Beginner Python for AI Engineering",
35
- ]
36
-
37
- DEFAULT_SELECTED_SOURCES_UI = [
38
- "Agentic AI Engineering",
39
- "Master AI For Work",
40
- "Full Stack AI Engineering",
41
- "Beginner Python for AI Engineering",
42
- "Transformers Docs",
43
- "PEFT Docs",
44
- "TRL Docs",
45
- "LlamaIndex Docs",
46
- "LangChain Docs",
47
- "OpenAI Cookbooks",
48
- ]
49
-
50
- AVAILABLE_SOURCES = [
51
- "transformers",
52
- "peft",
53
- "trl",
54
- "llama_index",
55
- "langchain",
56
- "openai_cookbooks",
57
- "full_stack_ai_engineering",
58
- "beginner_python_for_ai_engineering",
59
- "master_ai_for_work",
60
- "agentic_ai_engineering",
61
- ]
62
-
63
- SOURCE_UI_TO_KEY = {
64
- "Transformers Docs": "transformers",
65
- "PEFT Docs": "peft",
66
- "TRL Docs": "trl",
67
- "LlamaIndex Docs": "llama_index",
68
- "LangChain Docs": "langchain",
69
- "OpenAI Cookbooks": "openai_cookbooks",
70
- "Full Stack AI Engineering": "full_stack_ai_engineering",
71
- "Beginner Python for AI Engineering": "beginner_python_for_ai_engineering",
72
- "Master AI For Work": "master_ai_for_work",
73
- "Agentic AI Engineering": "agentic_ai_engineering",
74
- }
75
-
76
- COURSE_SOURCE_KEYS = frozenset(
77
- {
78
- "full_stack_ai_engineering",
79
- "beginner_python_for_ai_engineering",
80
- "master_ai_for_work",
81
- "agentic_ai_engineering",
82
- }
83
- )
84
-
85
- SOURCE_KEY_TO_LABEL = {value: key for key, value in SOURCE_UI_TO_KEY.items()}
86
- DEFAULT_SELECTED_SOURCE_KEYS = tuple(
87
- SOURCE_UI_TO_KEY[label] for label in DEFAULT_SELECTED_SOURCES_UI
88
- )
89
-
90
  CONCURRENCY_COUNT = int(os.getenv("CONCURRENCY_COUNT", 64))
91
  MONGODB_URI = os.getenv("MONGODB_URI")
92
 
@@ -122,6 +66,7 @@ __all__ = [
122
  "DEFAULT_SELECTED_SOURCES_UI",
123
  "CONCURRENCY_COUNT",
124
  "DEFAULT_MODEL_NAME",
 
125
  "DOCUMENT_DICT_PATH",
126
  "SOURCE_KEY_TO_LABEL",
127
  "SOURCE_UI_TO_KEY",
 
3
  import logfire
4
  from dotenv import load_dotenv
5
 
6
+ from data.scraping_scripts.source_registry import (
7
+ AVAILABLE_SOURCES,
8
+ AVAILABLE_SOURCES_UI,
9
+ COURSE_SOURCE_KEYS,
10
+ DEFAULT_SELECTED_SOURCE_KEYS,
11
+ DEFAULT_SELECTED_SOURCES_UI,
12
+ SOURCE_KEY_TO_LABEL,
13
+ SOURCE_UI_TO_KEY,
14
+ )
15
  from .utils import init_mongo_db
16
 
17
  load_dotenv(override=True)
 
23
  VECTOR_DB_DIR = "data/chroma-db-all_sources"
24
  VECTOR_COLLECTION_NAME = "chroma-db-all_sources"
25
  DOCUMENT_DICT_PATH = f"{VECTOR_DB_DIR}/document_dict_all_sources.pkl"
26
+ BM25_INDEX_PATH = f"{VECTOR_DB_DIR}/bm25_index_all_sources.pkl"
27
  DEFAULT_MODEL_NAME = "google-genai:gemini-3-flash-preview"
28
 
29
  AVAILABLE_MODELS: tuple[dict[str, str], ...] = (
 
31
  {"id": "anthropic:claude-haiku-4-5", "label": "Claude Haiku 4.5"},
32
  )
33
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
34
  CONCURRENCY_COUNT = int(os.getenv("CONCURRENCY_COUNT", 64))
35
  MONGODB_URI = os.getenv("MONGODB_URI")
36
 
 
66
  "DEFAULT_SELECTED_SOURCES_UI",
67
  "CONCURRENCY_COUNT",
68
  "DEFAULT_MODEL_NAME",
69
+ "BM25_INDEX_PATH",
70
  "DOCUMENT_DICT_PATH",
71
  "SOURCE_KEY_TO_LABEL",
72
  "SOURCE_UI_TO_KEY",
tests/test_chroma_rag.py ADDED
@@ -0,0 +1,201 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import asyncio
4
+ import pickle
5
+ import tempfile
6
+ import unittest
7
+ from pathlib import Path
8
+ from unittest.mock import patch
9
+
10
+ from data.scraping_scripts.add_context_to_nodes import process
11
+ from data.scraping_scripts.create_vector_stores import write_retrieval_artifacts
12
+ from llama_index.core import Document
13
+ from scripts.chroma_rag import (
14
+ BM25Index,
15
+ ChunkRecord,
16
+ build_chunk_records,
17
+ heading_aware_markdown_chunks,
18
+ reciprocal_rank_fusion,
19
+ load_bm25_index,
20
+ SearchResult,
21
+ )
22
+
23
+
24
+ class ChromaRagTestCase(unittest.TestCase):
25
+ def test_heading_aware_chunks_keep_code_blocks_intact(self) -> None:
26
+ code_lines = "\n".join(f"print({index})" for index in range(120))
27
+ markdown = f"""# Guide
28
+
29
+ ## Install
30
+
31
+ Use `pip install`.
32
+
33
+ ## Example
34
+
35
+ ```python
36
+ {code_lines}
37
+ ```
38
+
39
+ After the example.
40
+ """
41
+
42
+ chunks = heading_aware_markdown_chunks(
43
+ markdown,
44
+ title="Guide",
45
+ chunk_size=80,
46
+ )
47
+
48
+ code_chunks = [chunk for chunk in chunks if "print(0)" in chunk.text]
49
+ self.assertEqual(len(code_chunks), 1)
50
+ self.assertIn("print(119)", code_chunks[0].text)
51
+ self.assertIn("Example", code_chunks[0].heading_path)
52
+
53
+ def test_build_chunk_records_adds_heading_metadata(self) -> None:
54
+ records = build_chunk_records(
55
+ [
56
+ {
57
+ "doc_id": "doc-1",
58
+ "name": "Guide",
59
+ "url": "https://example.com/guide",
60
+ "source": "transformers",
61
+ "retrieve_doc": False,
62
+ "tokens": 1000,
63
+ "content": "# Guide\n\n## Install\n\nUse `AutoModel`.",
64
+ }
65
+ ]
66
+ )
67
+
68
+ self.assertEqual(records[0].metadata["heading_path"], "Guide")
69
+ self.assertEqual(records[1].metadata["heading_path"], "Guide > Install")
70
+ self.assertIn("source_version", records[0].metadata)
71
+
72
+ def test_bm25_search_finds_keywords_and_filters_sources(self) -> None:
73
+ records = [
74
+ ChunkRecord(
75
+ chunk_id="a",
76
+ doc_id="doc-a",
77
+ text="Use AutoModel.from_pretrained for model loading.",
78
+ metadata={"doc_id": "doc-a", "source": "transformers"},
79
+ ),
80
+ ChunkRecord(
81
+ chunk_id="b",
82
+ doc_id="doc-b",
83
+ text="Create a prompt template for chains.",
84
+ metadata={"doc_id": "doc-b", "source": "langchain"},
85
+ ),
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"]), [])
93
+
94
+ def test_retrieval_artifact_writer_persists_bm25_and_document_dict(self) -> None:
95
+ document_rows = [
96
+ {
97
+ "doc_id": "doc-1",
98
+ "name": "Transformers Loading",
99
+ "url": "https://example.com/loading",
100
+ "source": "transformers",
101
+ "retrieve_doc": False,
102
+ "tokens": 1200,
103
+ "content": "# Loading\n\n## AutoModel\n\nUse `AutoModel.from_pretrained`.",
104
+ }
105
+ ]
106
+
107
+ with tempfile.TemporaryDirectory(dir="/private/tmp") as temp_dir:
108
+ db_path = Path(temp_dir)
109
+ count = write_retrieval_artifacts(
110
+ config={
111
+ "document_dict_file": "document_dict_test.pkl",
112
+ "bm25_index_file": "bm25_index_test.pkl",
113
+ },
114
+ document_rows=document_rows,
115
+ db_path=str(db_path),
116
+ )
117
+
118
+ document_dict_path = db_path / "document_dict_test.pkl"
119
+ bm25_path = db_path / "bm25_index_test.pkl"
120
+
121
+ self.assertGreaterEqual(count, 1)
122
+ self.assertTrue(document_dict_path.exists())
123
+ self.assertTrue(bm25_path.exists())
124
+
125
+ with open(document_dict_path, "rb") as handle:
126
+ document_dict = pickle.load(handle)
127
+ self.assertEqual(document_dict["doc-1"]["name"], "Transformers Loading")
128
+
129
+ index = load_bm25_index(str(bm25_path))
130
+ self.assertIsNotNone(index)
131
+ assert index is not None
132
+ hits = index.search("AutoModel.from_pretrained")
133
+ self.assertEqual(hits[0][0].doc_id, "doc-1")
134
+ self.assertTrue(
135
+ any(record.metadata["heading_path"] for record in index.records)
136
+ )
137
+
138
+ def test_context_processing_uses_heading_chunks_and_raw_text_metadata(self) -> None:
139
+ async def fake_situate_context(_doc: str, chunk: str) -> str:
140
+ return f"Situated {chunk.splitlines()[0]}"
141
+
142
+ document = Document(
143
+ doc_id="doc-1",
144
+ text="# Guide\n\n## Setup\n\nUse `AutoModel.from_pretrained`.",
145
+ metadata={
146
+ "title": "Guide",
147
+ "url": "https://example.com/guide",
148
+ "tokens": 1000,
149
+ "retrieve_doc": False,
150
+ "source": "transformers",
151
+ },
152
+ )
153
+
154
+ with patch(
155
+ "data.scraping_scripts.add_context_to_nodes.situate_context",
156
+ fake_situate_context,
157
+ ):
158
+ records = asyncio.run(process([document], semaphore_limit=1))
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)
166
+ self.assertIn("Heading path: Guide > Setup", setup_record.text)
167
+ self.assertIn("Context: Situated", setup_record.text)
168
+
169
+ def test_rrf_prefers_overlap_across_ranked_lists(self) -> None:
170
+ dense_only = self._result("dense-only", 0.9, "dense")
171
+ overlap_dense = self._result("overlap", 0.7, "dense")
172
+ overlap_bm25 = self._result("overlap", 4.0, "bm25")
173
+ bm25_only = self._result("bm25-only", 5.0, "bm25")
174
+
175
+ fused = reciprocal_rank_fusion(
176
+ [[dense_only, overlap_dense], [bm25_only, overlap_bm25]],
177
+ top_k=4,
178
+ )
179
+
180
+ self.assertEqual(fused[0].chunk_id, "overlap")
181
+ self.assertEqual(fused[0].retrieval_method, "hybrid")
182
+
183
+ def _result(self, chunk_id: str, score: float, method: str) -> SearchResult:
184
+ return SearchResult(
185
+ chunk_id=chunk_id,
186
+ doc_id=chunk_id,
187
+ title=chunk_id,
188
+ url="",
189
+ source="test",
190
+ retrieve_doc=False,
191
+ tokens=10,
192
+ score=score,
193
+ content=chunk_id,
194
+ chunk_content=chunk_id,
195
+ heading_path="section",
196
+ retrieval_method=method,
197
+ )
198
+
199
+
200
+ if __name__ == "__main__":
201
+ unittest.main()
tests/test_process_md_files.py ADDED
@@ -0,0 +1,58 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from data.scraping_scripts.process_md_files import extract_title
2
+
3
+
4
+ def test_extract_title_prefers_frontmatter_title() -> None:
5
+ content = """---
6
+ title: "Dappier integration"
7
+ description: "Integrate with the Dappier retriever using LangChain Python."
8
+ ---
9
+
10
+ # DappierRetriever
11
+ """
12
+
13
+ assert extract_title(content) == "Dappier integration"
14
+
15
+
16
+ def test_extract_title_skips_frontmatter_delimiters() -> None:
17
+ content = """---
18
+ description: A page without an explicit title
19
+ ---
20
+
21
+ # Real page title
22
+ """
23
+
24
+ assert extract_title(content) == "Real page title"
25
+
26
+
27
+ def test_extract_title_uses_sidebar_title_when_title_missing() -> None:
28
+ content = """---
29
+ sidebarTitle: Overview
30
+ description: A page without an explicit title
31
+ ---
32
+
33
+ # Longer body heading
34
+ """
35
+
36
+ assert extract_title(content) == "Overview"
37
+
38
+
39
+ def test_extract_title_ignores_headings_inside_code_fences() -> None:
40
+ content = """```python
41
+ # Not a page title
42
+ ```
43
+
44
+ # Real page title
45
+ """
46
+
47
+ assert extract_title(content) == "Real page title"
48
+
49
+
50
+ def test_extract_title_keeps_first_line_fallback() -> None:
51
+ content = """---
52
+ description: A page without heading syntax
53
+ ---
54
+
55
+ First paragraph fallback
56
+ """
57
+
58
+ assert extract_title(content) == "First paragraph fallback"
tests/test_retire_source_workflow.py ADDED
@@ -0,0 +1,54 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import ast
2
+
3
+ from data.scraping_scripts.retire_source_workflow import (
4
+ remove_source_from_registry_text,
5
+ )
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"},
13
+ },
14
+ "langchain": {
15
+ "output_file": "data/langchain_data.jsonl",
16
+ },
17
+ }
18
+
19
+ DOC_SOURCE_KEYS = (
20
+ "openai_cookbooks",
21
+ "langchain",
22
+ )
23
+ GITHUB_SOURCE_KEYS = (
24
+ "openai_cookbooks",
25
+ "langchain",
26
+ )
27
+ SOURCE_KEY_TO_LABEL = {
28
+ "openai_cookbooks": "OpenAI Cookbooks",
29
+ "langchain": "LangChain Docs",
30
+ }
31
+ UI_SOURCE_KEYS = (
32
+ "openai_cookbooks",
33
+ "langchain",
34
+ )
35
+ '''
36
+
37
+ updated, changed = remove_source_from_registry_text(
38
+ registry,
39
+ "openai_cookbooks",
40
+ )
41
+
42
+ assert changed is True
43
+ assert "openai_cookbooks" not in updated
44
+ assert '"langchain"' in updated
45
+ ast.parse(updated)
46
+
47
+
48
+ def test_remove_source_from_registry_text_ignores_missing_source():
49
+ registry = 'SOURCE_CONFIGS = {"langchain": {"output_file": "x"}}\n'
50
+
51
+ updated, changed = remove_source_from_registry_text(registry, "missing")
52
+
53
+ assert changed is False
54
+ assert updated == registry