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 +24 -14
- data/scraping_scripts/add_context_to_nodes.py +214 -67
- data/scraping_scripts/add_course_workflow.py +46 -115
- data/scraping_scripts/capture_source_versions.py +50 -8
- data/scraping_scripts/contextual_node_pruning.py +67 -0
- data/scraping_scripts/create_vector_stores.py +237 -85
- data/scraping_scripts/github_to_markdown_ai_docs.py +60 -4
- data/scraping_scripts/llms_txt_to_markdown_docs.py +194 -0
- data/scraping_scripts/process_md_files.py +176 -167
- data/scraping_scripts/retire_source_workflow.py +114 -4
- data/scraping_scripts/source_registry.py +358 -0
- data/scraping_scripts/update_docs_workflow.py +85 -39
- data/scraping_scripts/upload_data_to_hf.py +3 -19
- frontend/lib/doc-metadata.ts +20 -0
- scripts/chat_service.py +2 -0
- scripts/chroma_rag.py +909 -72
- scripts/setup.py +11 -66
- tests/test_chroma_rag.py +201 -0
- tests/test_process_md_files.py +58 -0
- tests/test_retire_source_workflow.py +54 -0
data/scraping_scripts/README.md
CHANGED
|
@@ -4,11 +4,18 @@
|
|
| 4 |
|
| 5 |
Make sure you have the required environment variables set:
|
| 6 |
|
| 7 |
-
- `
|
| 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
|
| 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 --
|
| 74 |
```
|
| 75 |
|
| 76 |
example:
|
| 77 |
|
| 78 |
```bash
|
| 79 |
-
uv run -m data.scraping_scripts.add_course_workflow --
|
| 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.
|
| 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 --
|
| 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.
|
| 176 |
-
|
| 177 |
-
|
|
|
|
|
|
|
| 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 |
-
|
| 186 |
-
|
| 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 |
-
-
|
| 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,
|
| 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 |
-
|
| 65 |
-
|
| 66 |
-
|
| 67 |
-
|
| 68 |
-
|
| 69 |
-
|
| 70 |
-
|
| 71 |
-
|
| 72 |
-
|
| 73 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 74 |
|
|
|
|
| 75 |
|
| 76 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 92 |
"""
|
| 93 |
)
|
| 94 |
|
| 95 |
content = template.render(doc=doc, chunk=chunk)
|
| 96 |
-
|
| 97 |
-
|
| 98 |
-
|
| 99 |
-
|
| 100 |
-
|
| 101 |
-
|
| 102 |
-
|
| 103 |
-
|
| 104 |
-
|
| 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 |
-
|
| 125 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 126 |
doc: Document = document_dict[doc_id]
|
| 127 |
|
| 128 |
-
if doc.metadata["tokens"] >
|
| 129 |
# Tokenize the document text
|
| 130 |
-
encoding = tiktoken.
|
| 131 |
tokens = encoding.encode(doc.get_content())
|
| 132 |
|
| 133 |
# Trim to 120,000 tokens
|
| 134 |
-
trimmed_tokens = tokens[:
|
| 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"] =
|
| 142 |
-
|
| 143 |
-
context
|
| 144 |
-
|
| 145 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 146 |
|
| 147 |
|
| 148 |
async def process(
|
| 149 |
-
documents: List[Document], semaphore_limit: int =
|
| 150 |
) -> List[ChunkRecord]:
|
| 151 |
|
| 152 |
-
|
| 153 |
-
|
| 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(
|
| 164 |
async with semaphore:
|
| 165 |
-
result = await process_chunk(
|
| 166 |
await asyncio.sleep(0.1)
|
| 167 |
return result
|
| 168 |
|
| 169 |
-
tasks = [process_with_semaphore(
|
| 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
|
| 10 |
-
(this naturally drops any retired sources no longer in
|
| 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 |
-
|
| 75 |
-
|
| 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
|
| 210 |
|
| 211 |
Unlike the previous append-style merge, this drops any source whose entry
|
| 212 |
-
has been removed from
|
| 213 |
-
|
| 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 |
-
"""
|
| 419 |
-
|
| 420 |
-
|
| 421 |
-
|
| 422 |
-
|
| 423 |
-
|
| 424 |
-
|
| 425 |
-
logger.error(f"Course {course_name} not found in SOURCE_CONFIGS")
|
| 426 |
return
|
| 427 |
|
| 428 |
-
|
| 429 |
-
|
| 430 |
-
|
| 431 |
-
|
| 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
|
| 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
|
| 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
|
| 11 |
-
|
| 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
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
| 119 |
print(f" ! skipping unknown source: {source}", file=sys.stderr)
|
| 120 |
continue
|
| 121 |
-
versions[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(
|
| 137 |
-
default=list(
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
"
|
| 165 |
-
len(
|
|
|
|
| 166 |
source,
|
| 167 |
-
|
| 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 |
-
|
| 185 |
-
|
| 186 |
-
|
| 187 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 188 |
)
|
| 189 |
|
| 190 |
print(
|
| 191 |
-
f"Indexed {len(chunk_records)} chunks
|
|
|
|
| 192 |
)
|
| 193 |
|
| 194 |
|
| 195 |
-
def main(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 196 |
for source in sources:
|
| 197 |
if source not in SOURCE_CONFIGS:
|
| 198 |
print(f"Unknown source: {source}")
|
| 199 |
continue
|
| 200 |
-
process_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(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
|
|
|
| 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
|
| 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 |
-
|
| 200 |
-
|
| 201 |
-
|
| 202 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 203 |
|
| 204 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 278 |
-
|
|
|
|
| 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(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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.
|
| 9 |
-
4. the
|
|
|
|
| 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.
|
| 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
|
| 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
|
| 6 |
-
1. Download documentation from GitHub
|
| 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 |
-
|
| 68 |
-
|
| 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
|
| 133 |
-
|
| 134 |
-
|
| 135 |
-
|
| 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=
|
| 371 |
-
default=
|
| 372 |
-
help="
|
| 373 |
)
|
| 374 |
parser.add_argument(
|
| 375 |
-
"--skip-download", action="store_true", help="Skip downloading
|
| 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 |
-
#
|
| 407 |
-
|
|
|
|
|
|
|
| 408 |
|
| 409 |
# Execute the workflow steps
|
| 410 |
if not args.skip_download:
|
| 411 |
-
|
| 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 |
-
|
| 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 =
|
| 107 |
document["content"],
|
|
|
|
| 108 |
chunk_size=chunk_size,
|
| 109 |
chunk_overlap=chunk_overlap,
|
| 110 |
)
|
| 111 |
-
|
|
|
|
|
|
|
|
|
|
| 112 |
metadata = {
|
| 113 |
"doc_id": document["doc_id"],
|
| 114 |
"title": document["name"],
|
| 115 |
"url": document["url"],
|
| 116 |
-
"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=
|
| 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 =
|
|
|
|
|
|
|
|
|
|
| 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
|
| 290 |
-
|
| 291 |
-
|
| 292 |
-
|
| 293 |
-
|
| 294 |
-
|
| 295 |
-
|
| 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 |
-
|
| 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 |
-
|
| 485 |
-
|
|
|
|
| 486 |
)
|
| 487 |
)
|
|
|
|
| 488 |
|
| 489 |
-
|
| 490 |
-
|
| 491 |
-
|
| 492 |
-
|
| 493 |
-
|
| 494 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|