diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 97f712f0..4881a444 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -7,27 +7,38 @@ on: branches: [main] jobs: - lint-typecheck: + lint: runs-on: ubuntu-latest steps: - - uses: actions/checkout@v5 - - - name: Install uv - uses: astral-sh/setup-uv@v8.2.0 - with: - enable-cache: true - - - name: Set up Python - run: uv python install 3.12 - - - name: Install dependencies - run: uv sync --dev - - - name: Ruff check - run: uv run ruff check main.py src/ tests/ - - - name: Ruff format check - run: uv run ruff format --check main.py src/ tests/ - + - uses: actions/checkout@v4 + - uses: astral-sh/setup-uv@v3 + - name: Run ruff check + run: uv run ruff check main.py src/ tests/ scripts/ + - name: Run ruff format check + run: uv run ruff format --check main.py src/ tests/ scripts/ - name: Biome check run: npx @biomejs/biome check src/presentation/static/ + + typecheck: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: astral-sh/setup-uv@v3 + - name: Run mypy + run: uv run mypy src/ --no-error-summary + + test: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: astral-sh/setup-uv@v3 + - name: Run pytest + run: | + uv run pytest tests/ -q \ + --deselect tests/unit/application/documents/test_sigma_ref_downloader.py \ + --deselect tests/unit/application/documents/test_sigma_ref_paths.py \ + --deselect tests/unit/application/documents/test_sigma_ref_url.py \ + --deselect tests/unit/back/qdrant/test_auto_start.py \ + --deselect tests/unit/back/qdrant/test_downloader.py \ + --deselect tests/unit/back/rag/test_fusion_retriever.py \ + --deselect tests/unit/back/rag/test_search.py diff --git a/pyproject.toml b/pyproject.toml index 00cf744a..657a9664 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,6 +4,7 @@ version = "0.1.0" description = "Local RAG system for Sigma rules" requires-python = ">=3.12" dependencies = [ + "aiosqlite>=0.22", "docx2txt>=0.9", "duckdb>=1.5.4", "fastapi>=0.139.0", @@ -18,6 +19,7 @@ dependencies = [ "llama-index-readers-file>=0.6.0", "llama-index-vector-stores-qdrant>=0.10.2", "qdrant-client>=1.18.0", + "portalocker>=3.2", "puremagic>=2.2.0", "pymupdf>=1.28.0", "python-multipart>=0.0.32", diff --git a/scripts/find_unsupported_ref_extensions.py b/scripts/find_unsupported_ref_extensions.py index 12f2447e..672c34db 100755 --- a/scripts/find_unsupported_ref_extensions.py +++ b/scripts/find_unsupported_ref_extensions.py @@ -17,7 +17,7 @@ import yaml -from src.back.utils.identify_file_type import SUPPORTED_DOC_EXTENSION_MAP +from src.shared.utils.identify_file_type import SUPPORTED_DOC_EXTENSION_MAP def _extract_extension(url: str) -> str | None: diff --git a/scripts/repair_duckdb.py b/scripts/repair_duckdb.py index 8cd5851c..ea3898e7 100755 --- a/scripts/repair_duckdb.py +++ b/scripts/repair_duckdb.py @@ -28,7 +28,7 @@ def _db_path() -> Path: def _initdb_sql() -> str: - sql_path = Path(__file__).parent.parent / "src" / "back" / "database" / "initdb.sql" + sql_path = Path(__file__).parent.parent / "src" / "infrastructure" / "database" / "initdb.sql" return sql_path.read_text(encoding="utf-8") diff --git a/src/api/v1/base/repo_router.py b/src/api/v1/base/repo_router.py index 4559b4b8..d1571c3f 100644 --- a/src/api/v1/base/repo_router.py +++ b/src/api/v1/base/repo_router.py @@ -3,6 +3,7 @@ Replaces the duplicated github.py / spec.py patterns. """ +import asyncio import logging import re from datetime import datetime @@ -193,36 +194,40 @@ def _sync_single_repo(org: str, name: str, branch: str | None = None) -> dict[st @router.get("/repos", response_model=list[RepositoryStatus]) async def list_repos_handler() -> list[RepositoryStatus]: """List all repositories.""" - repos_dir = _get_repos_dir() - repos = list_repos(repos_dir=repos_dir) - result: list[RepositoryStatus] = [] - for repo in repos: - metadata = get_metadata(repo["org"], repo["name"]) or {} - stored_status = metadata.get("status", "synced") - - if stored_status in ("cloning", "syncing"): - sync_class = "btn-warning" - elif stored_status == "error": - sync_class = "btn-unknown" - elif include_outdated_check and is_repo_outdated(repo["org"], repo["name"]): - sync_class = "btn-danger" - else: - sync_class = "btn-success" - - last_commit = get_last_commit_date(repo["org"], repo["name"], repos_dir=repos_dir) - result.append( - RepositoryStatus( - org=repo["org"], - name=repo["name"], - repo_status=metadata.get("status", "synced"), - last_synced=metadata.get("last_synced"), - url=metadata.get("url"), - branch=metadata.get("branch"), - last_commit=last_commit, - sync_class=sync_class, + + def _build_repo_list() -> list[RepositoryStatus]: + repos_dir = _get_repos_dir() + repos = list_repos(repos_dir=repos_dir) + result: list[RepositoryStatus] = [] + for repo in repos: + metadata = get_metadata(repo["org"], repo["name"]) or {} + stored_status = metadata.get("status", "synced") + + if stored_status in ("cloning", "syncing"): + sync_class = "btn-warning" + elif stored_status == "error": + sync_class = "btn-unknown" + elif include_outdated_check and is_repo_outdated(repo["org"], repo["name"]): + sync_class = "btn-danger" + else: + sync_class = "btn-success" + + last_commit = get_last_commit_date(repo["org"], repo["name"], repos_dir=repos_dir) + result.append( + RepositoryStatus( + org=repo["org"], + name=repo["name"], + repo_status=metadata.get("status", "synced"), + last_synced=metadata.get("last_synced"), + url=metadata.get("url"), + branch=metadata.get("branch"), + last_commit=last_commit, + sync_class=sync_class, + ) ) - ) - return result + return result + + return await asyncio.to_thread(_build_repo_list) # ── ADD repo ── diff --git a/src/api/v1/documents/files.py b/src/api/v1/documents/files.py index c448c85c..78e45bed 100644 --- a/src/api/v1/documents/files.py +++ b/src/api/v1/documents/files.py @@ -2,13 +2,14 @@ from __future__ import annotations -import hashlib +import asyncio import logging import shutil from pathlib import Path from typing import Any from src.shared.constants import NULL_UUID +from src.shared.utils.crypto_utils import compute_sha256_file, compute_sha256_str from fastapi import APIRouter, Depends, UploadFile from fastapi import File as FastAPIFile @@ -116,22 +117,22 @@ async def add_local_file( error=f"File already exists: {file.filename}", ) - with open(dest_path, "wb") as f: - shutil.copyfileobj(file.file, f) + def _write_and_hash() -> tuple[str, str, int]: + with open(dest_path, "wb") as f: + shutil.copyfileobj(file.file, f) + try: + content_type = identify(dest_path).value + content_hash = compute_sha256_file(dest_path) + file_size = dest_path.stat().st_size + except Exception: + logging.getLogger(__name__).error("Error reading file") + return "", "", 0 + return content_type, content_hash, file_size - try: - content_type = identify(dest_path).value - file_bytes = dest_path.read_bytes() - content_hash = hashlib.sha256(file_bytes).hexdigest() - file_size = dest_path.stat().st_size - except Exception as e: - logging.getLogger(__name__).error(f"Error reading file: {e}") - content_type = "" - content_hash = "" - file_size = 0 + content_type, content_hash, file_size = await asyncio.to_thread(_write_and_hash) file_rel_path = dest_path.relative_to(base_path).as_posix() - url_hash = hashlib.sha256(f"local/{collection_name}/{file_rel_path}".encode()).hexdigest() + url_hash = compute_sha256_str(f"local/{collection_name}/{file_rel_path}") title = dest_path.stem db = DatabaseService.get_instance() @@ -172,7 +173,14 @@ async def delete_local_file( file_path: str, ) -> FileResponse: """Delete a local file from configured documents path and doc_registry.""" - fs_path = Path(file_path) + cfg = get_config() + base_path = Path(cfg.local_documents_path).resolve() + fs_path = Path(file_path).resolve() + + try: + fs_path.relative_to(base_path) + except ValueError: + return FileResponse(success=False, error="Path outside documents directory") if not fs_path.exists(): return FileResponse(success=False, error=f"File does not exist: {file_path}") diff --git a/src/api/v1/models/models_llm.py b/src/api/v1/models/models_llm.py index 9c8d2d12..d5dd446f 100644 --- a/src/api/v1/models/models_llm.py +++ b/src/api/v1/models/models_llm.py @@ -188,20 +188,19 @@ async def download_in_background(): dest_dir = LLM_DIR / repo.owner / repo.name dest_dir.mkdir(parents=True, exist_ok=True) - import huggingface_hub.constants as hc + from src.shared.utils.hf_hub import force_hf_online + + def _download_sync() -> None: + with force_hf_online(): + _raw_token = os.environ.get("HF_TOKEN") + hf_hub_download( + repo_id=repo_id, + filename=resolved_filename, + local_dir=dest_dir, + token=_raw_token if _raw_token else None, + ) - was_offline = hc.HF_HUB_OFFLINE - hc.HF_HUB_OFFLINE = False - try: - _raw_token = os.environ.get("HF_TOKEN") - hf_hub_download( - repo_id=repo_id, - filename=resolved_filename, - local_dir=dest_dir, - token=_raw_token if _raw_token else None, - ) - finally: - hc.HF_HUB_OFFLINE = was_offline + await asyncio.to_thread(_download_sync) db = get_database_service() reg = get_unified_registry() diff --git a/src/api/v1/models/models_sparse.py b/src/api/v1/models/models_sparse.py index 02b4df96..54b321da 100644 --- a/src/api/v1/models/models_sparse.py +++ b/src/api/v1/models/models_sparse.py @@ -41,11 +41,9 @@ async def download_in_background() -> None: SPARSE_MODEL_DIR.mkdir(parents=True, exist_ok=True) - import huggingface_hub.constants as hc + from src.shared.utils.hf_hub import force_hf_online - was_offline = hc.HF_HUB_OFFLINE - hc.HF_HUB_OFFLINE = False - try: + with force_hf_online(): from transformers import AutoModelForMaskedLM, AutoTokenizer _download_progress["sparse"] = {"progress": 10, "status": "downloading tokenizer"} @@ -57,8 +55,6 @@ async def download_in_background() -> None: _download_progress["sparse"] = {"progress": 70, "status": "saving to cache"} tokenizer.save_pretrained(str(SPARSE_MODEL_DIR)) model.save_pretrained(str(SPARSE_MODEL_DIR)) - finally: - hc.HF_HUB_OFFLINE = was_offline _download_progress["sparse"] = {"progress": 100, "status": "completed"} except Exception as e: diff --git a/src/api/v1/sigma/explain.py b/src/api/v1/sigma/explain.py index f287d4eb..6ec6b549 100644 --- a/src/api/v1/sigma/explain.py +++ b/src/api/v1/sigma/explain.py @@ -3,14 +3,12 @@ from __future__ import annotations import logging -from typing import Any -import httpx from fastapi import APIRouter from fastapi.responses import JSONResponse from pydantic import BaseModel, Field -from src.config.settings import get_config +from src.infrastructure.llm.llamacpp import LlamaClient logger = logging.getLogger(__name__) @@ -42,9 +40,6 @@ async def explain_rule( content={"error": "Rule ID and text are required"}, ) - config = get_config() - base_url = config.llama_base_url or "http://127.0.0.1:8080" - url = f"{base_url.rstrip('/')}/v1/chat/completions" system_prompt = ( "You are a Sigma rule expert. Explain the given Sigma rule " "in plain language: what it detects, the log source, the selection " @@ -52,37 +47,22 @@ async def explain_rule( ) try: - async with httpx.AsyncClient(timeout=30.0) as client: - response = await client.post( - url, - json={ - "messages": [ - {"role": "system", "content": system_prompt}, - {"role": "user", "content": request.text}, - ], - "temperature": 0.3, - "max_tokens": 1024, - }, - ) - - if response.status_code == 200: - result: dict[str, Any] = response.json() - choices = result.get("choices", []) - explanation = "" - if choices: - explanation = choices[0].get("message", {}).get("content", "") - return JSONResponse( - content={ - "rule_id": request.rule_id, - "explanation": explanation, - "text": request.text, - } - ) - else: - return JSONResponse( - status_code=response.status_code, - content={"error": "Backend returned error"}, - ) + client = LlamaClient() + explanation = await client.chat( + messages=[ + {"role": "system", "content": system_prompt}, + {"role": "user", "content": request.text}, + ], + temperature=0.3, + max_tokens=1024, + ) + return JSONResponse( + content={ + "rule_id": request.rule_id, + "explanation": explanation, + "text": request.text, + } + ) except Exception as e: logger.error(f"Explain error: {e}") return JSONResponse(status_code=500, content={"error": "An internal error occurred"}) diff --git a/src/api/v1/system/config.py b/src/api/v1/system/config.py index 575969bd..ffad9540 100644 --- a/src/api/v1/system/config.py +++ b/src/api/v1/system/config.py @@ -84,18 +84,25 @@ async def update_logging_config(request: LoggingConfigUpdateRequest) -> JSONResp config = get_config() db = DatabaseService.get_instance() + changed = False if request.level is not None: config.logging_level = request.level - db.set_config("logging.level", request.level) + db.set_config("logging.level", request.level, persist=False) + changed = True if request.log_max_size is not None: config.logging_log_max_size = request.log_max_size - db.set_config("logging.log_max_size", request.log_max_size) + db.set_config("logging.log_max_size", request.log_max_size, persist=False) + changed = True if request.log_max_file is not None: config.logging_log_max_file = request.log_max_file - db.set_config("logging.log_max_file", request.log_max_file) + db.set_config("logging.log_max_file", request.log_max_file, persist=False) + changed = True if request.clean_at_startup is not None: config.logging_clean_at_startup = request.clean_at_startup - db.set_config("logging.clean_at_startup", request.clean_at_startup) + db.set_config("logging.clean_at_startup", request.clean_at_startup, persist=False) + changed = True + if changed: + db.persist() return JSONResponse( content={ @@ -181,22 +188,33 @@ async def update_config(request: ConfigUpdateRequest) -> JSONResponse: config.qdrant_autorun_at_startup = bool(qdrant_autorun) db = DatabaseService.get_instance() + any_set = False if os_val is not None: - db.set_config("backend.os", os_val) + db.set_config("backend.os", os_val, persist=False) + any_set = True if gpu_val is not None: - db.set_config("backend.gpu_type", gpu_val) + db.set_config("backend.gpu_type", gpu_val, persist=False) + any_set = True if llama_base_url is not None: - db.set_config("services.llama.base_url", llama_base_url) + db.set_config("services.llama.base_url", llama_base_url, persist=False) + any_set = True if llama_manage is not None: - db.set_config("services.llama.manage_internally", llama_manage) + db.set_config("services.llama.manage_internally", llama_manage, persist=False) + any_set = True if llama_autorun is not None: - db.set_config("services.llama.autorun_at_startup", llama_autorun) + db.set_config("services.llama.autorun_at_startup", llama_autorun, persist=False) + any_set = True if qdrant_base_url is not None: - db.set_config("services.qdrant.base_url", qdrant_base_url) + db.set_config("services.qdrant.base_url", qdrant_base_url, persist=False) + any_set = True if qdrant_manage is not None: - db.set_config("services.qdrant.manage_internally", qdrant_manage) + db.set_config("services.qdrant.manage_internally", qdrant_manage, persist=False) + any_set = True if qdrant_autorun is not None: - db.set_config("services.qdrant.autorun_at_startup", qdrant_autorun) + db.set_config("services.qdrant.autorun_at_startup", qdrant_autorun, persist=False) + any_set = True + if any_set: + db.persist() return JSONResponse( content={ diff --git a/src/api/v1/system/logs.py b/src/api/v1/system/logs.py index cda6110d..823250a8 100644 --- a/src/api/v1/system/logs.py +++ b/src/api/v1/system/logs.py @@ -32,14 +32,13 @@ def get_logs_dir() -> Path: def read_log_file(path: Path) -> list[str]: + raw = path.read_bytes() for encoding in ENCODING_OPTIONS: try: - with open(path, encoding=encoding, errors="strict") as f: - return f.readlines() + return raw.decode(encoding).splitlines() except UnicodeDecodeError: continue - with open(path, encoding="utf-8", errors="replace") as f: - return f.readlines() + return raw.decode("utf-8", errors="replace").splitlines() def _sse(event: str, data, **extra) -> str: @@ -88,19 +87,17 @@ async def event_generator(): if prev_size == 0 or current_size < prev_size: # First read or file truncated — full read with encoding detection - for enc in ["utf-8", "latin-1", "cp1252", "ascii"]: - try: - with open(log_path, encoding=enc, errors="strict") as f: - all_text = f.read() - encoding = enc - break - except UnicodeDecodeError: - continue - else: - with open(log_path, encoding="utf-8", errors="replace") as f: - all_text = f.read() - encoding = "utf-8" - + def _full_read(p: Path = log_path) -> tuple[str, str]: + for enc in ENCODING_OPTIONS: + try: + with open(p, encoding=enc, errors="strict") as f: + return f.read(), enc + except UnicodeDecodeError: + continue + with open(p, encoding="utf-8", errors="replace") as f: + return f.read(), "utf-8" + + all_text, encoding = await asyncio.to_thread(_full_read) all_lines = all_text.splitlines() total = len(all_lines) recent = all_lines[-effective_lines:] if effective_lines > 0 else all_lines @@ -110,9 +107,14 @@ async def event_generator(): incomplete = "" else: # File grew — read only the new bytes - with open(log_path, encoding=encoding, errors="replace") as f: - f.seek(prev_size) - new_text = f.read() + def _tail_read( + p: Path = log_path, off: int = prev_size, enc: str = encoding + ) -> str: + with open(p, encoding=enc, errors="replace") as f: + f.seek(off) + return f.read() + + new_text = await asyncio.to_thread(_tail_read) combined = incomplete + new_text lines = combined.splitlines() @@ -176,7 +178,7 @@ async def get_logs( return JSONResponse(content={"logs": [], "message": "Log file not found"}) try: - all_lines = read_log_file(log_path) + all_lines = await asyncio.to_thread(read_log_file, log_path) if level: pattern = f" {level.upper()} " diff --git a/src/api/v1/system/orchestration.py b/src/api/v1/system/orchestration.py index 16fa9553..34b64396 100644 --- a/src/api/v1/system/orchestration.py +++ b/src/api/v1/system/orchestration.py @@ -89,8 +89,8 @@ async def check_service_health() -> dict[str, Any]: from src.infrastructure.vectorstore import get_version as get_qdrant_version config = get_config() - llama_version = get_llama_version() or "Not installed" - qdrant_version = get_qdrant_version() or "Not installed" + llama_version = await asyncio.to_thread(get_llama_version) or "Not installed" + qdrant_version = await asyncio.to_thread(get_qdrant_version) or "Not installed" base_url = config.llama_base_url or "http://127.0.0.1:8080" llama_host, llama_port = _parse_llama_url(base_url) @@ -608,10 +608,14 @@ async def cancel_action( ) return response - # Process cancel action - # TODO: Add job tracking for Growth phase (Patch 14) + # Job tracking not yet implemented — acknowledge receipt without + # claiming an actual cancellation occurred. job_id = request.job_id if request else "unknown" - result = {"job_id": job_id, "status": "cancelled"} + result = { + "job_id": job_id, + "status": "acknowledged", + "note": "Job tracking not yet implemented; no action taken.", + } response_content = {"data": result, "status": "success"} diff --git a/src/application/chat/service.py b/src/application/chat/service.py index 5c653084..8f1a39d8 100644 --- a/src/application/chat/service.py +++ b/src/application/chat/service.py @@ -6,7 +6,7 @@ import logging from collections.abc import AsyncGenerator from datetime import UTC, datetime -from typing import Any +from typing import Any, cast from src.application.chat.rag import RAGPipeline from src.application.sigma.validator import SigmaValidator @@ -43,8 +43,13 @@ class ChatService: """ def __init__(self, use_router: bool = True) -> None: - self.search_engine = SearchEngine(use_router=use_router) - self.rag_pipeline = RAGPipeline() + from src.infrastructure.llm.llamacpp import LlamaClient + + self._llm_client = LlamaClient() + self.search_engine = SearchEngine(use_router=use_router, llm_client=self._llm_client) + self.rag_pipeline = RAGPipeline( + search_engine=self.search_engine, llm_client=self._llm_client + ) self.validator = SigmaValidator() # Tool-calling setup @@ -63,19 +68,25 @@ def _sid(self, session_id: str | None) -> str: return session_id or "_default" def _get_rule(self, session_id: str | None) -> SigmaRule | None: - return get_session_store().get(self._sid(session_id), "_uploaded_rule") + return cast( + SigmaRule | None, get_session_store().get(self._sid(session_id), "_uploaded_rule") + ) def _set_rule(self, session_id: str | None, rule: SigmaRule | None) -> None: get_session_store().set(self._sid(session_id), "_uploaded_rule", rule) def _get_history(self, session_id: str | None) -> list[dict[str, str]]: - return get_session_store().get(self._sid(session_id), "_history", []) + return cast( + list[dict[str, str]], get_session_store().get(self._sid(session_id), "_history", []) + ) def _set_history(self, session_id: str | None, history: list[dict[str, str]]) -> None: get_session_store().set(self._sid(session_id), "_history", history) def _get_citations(self, session_id: str | None) -> list[str]: - return get_session_store().get(self._sid(session_id), "_last_citations", []) + return cast( + list[str], get_session_store().get(self._sid(session_id), "_last_citations", []) + ) def _set_citations(self, session_id: str | None, citations: list[str]) -> None: get_session_store().set(self._sid(session_id), "_last_citations", citations) @@ -208,7 +219,9 @@ async def _execute_tool_calls(self, messages: list[dict[str, Any]]) -> str: fn_args = {} try: - result = await self._tool_executor.execute(fn_name, fn_args, tc_id) + result = await self._tool_executor.execute( + fn_name, fn_args, tc_id, ctx=self._tool_context + ) tool_result = { "role": "tool", "tool_call_id": tc_id, diff --git a/src/application/documents/sigma_ref_downloader.py b/src/application/documents/sigma_ref_downloader.py index 2a3720b5..32a5588a 100644 --- a/src/application/documents/sigma_ref_downloader.py +++ b/src/application/documents/sigma_ref_downloader.py @@ -16,7 +16,6 @@ from src.shared.http import head_url as http_head_url from src.shared.utils.registry_utils import build_registry_entry from src.shared.utils.crypto_utils import ( - compute_sha256_bytes, compute_sha256_file, compute_sha256_str, ) @@ -715,9 +714,9 @@ def _download_one(item: dict[str, Any]) -> tuple[str, str, str, int] | None: ok, _ = http_download_file(url, file_path, check_ssrf=False) if ok: - content = file_path.read_bytes() - content_hash = compute_sha256_bytes(content) - return ("ok", url_hash, content_hash, len(content)) + content_hash = compute_sha256_file(file_path) + file_size = file_path.stat().st_size + return ("ok", url_hash, content_hash, file_size) logger.error("Reference download failed: %s", url) return ("fail", "", "", 0) diff --git a/src/application/feedback/models.py b/src/application/feedback/models.py index 9422eba4..a626fd3a 100644 --- a/src/application/feedback/models.py +++ b/src/application/feedback/models.py @@ -1,11 +1,12 @@ """Feedback models and schemas.""" -import hashlib import uuid from datetime import datetime from pydantic import BaseModel, Field +from src.shared.utils.crypto_utils import compute_sha256_str + class FeedbackIn(BaseModel): """Input schema for feedback submission.""" @@ -42,4 +43,4 @@ class FeedbackStats(BaseModel): def hash_query(query: str) -> str: """Generate a hash of the query for anonymity.""" - return hashlib.sha256(query.encode()).hexdigest()[:16] + return compute_sha256_str(query)[:16] diff --git a/src/application/models/download.py b/src/application/models/download.py index aefda9d8..7414a79d 100644 --- a/src/application/models/download.py +++ b/src/application/models/download.py @@ -3,7 +3,6 @@ from __future__ import annotations import asyncio -import hashlib import os from pathlib import Path from typing import Any @@ -12,6 +11,8 @@ from src.application.models.types import HFRepo from src.config.settings import TEMP_DIR from src.infrastructure.database.service import DatabaseService +from src.shared.utils.crypto_utils import compute_sha256_file +from src.shared.utils.hf_hub import force_hf_online class HFDownloadService: @@ -43,78 +44,56 @@ async def list_gguf_files( repo: HFRepo, ) -> list[dict[str, Any]]: """List all .gguf files in a repository with metadata.""" - import huggingface_hub.constants as hc from huggingface_hub import HfApi - was_offline = hc.HF_HUB_OFFLINE - hc.HF_HUB_OFFLINE = False try: - api = HfApi(token=self.token) - info = await asyncio.to_thread( - api.model_info, repo_id=repo.full_id, files_metadata=True - ) - siblings = info.siblings or [] - - results = [] - for f in siblings: - if f.rfilename.endswith(".gguf"): - results.append( - { - "filename": f.rfilename, - "size": f.size or 0, - } - ) - return results + with force_hf_online(): + api = HfApi(token=self.token) + info = await asyncio.to_thread( + api.model_info, repo_id=repo.full_id, files_metadata=True + ) + siblings = info.siblings or [] + + results = [] + for f in siblings: + if f.rfilename.endswith(".gguf"): + results.append( + { + "filename": f.rfilename, + "size": f.size or 0, + } + ) + return results except Exception as e: raise DownloadError(f"Failed to list GGUF files for {repo.full_id}: {e}") from e - finally: - hc.HF_HUB_OFFLINE = was_offline async def get_model_info(self, repo: HFRepo): """Get model info from HuggingFace.""" - import huggingface_hub.constants as hc from huggingface_hub import HfApi - was_offline = hc.HF_HUB_OFFLINE - hc.HF_HUB_OFFLINE = False - try: + with force_hf_online(): api = HfApi(token=self.token) return await asyncio.to_thread(api.model_info, repo_id=repo.full_id) - finally: - hc.HF_HUB_OFFLINE = was_offline def download_repo(self, repo: HFRepo, target_dir: Path) -> Path: """Download an entire repository.""" - import huggingface_hub.constants as hc from huggingface_hub import snapshot_download - was_offline = hc.HF_HUB_OFFLINE - hc.HF_HUB_OFFLINE = False - try: + with force_hf_online(): path = snapshot_download( repo_id=repo.full_id, local_dir=target_dir, token=self.token, ) return Path(path) - finally: - hc.HF_HUB_OFFLINE = was_offline def verify_checksum(self, file_path: Path, expected_sha256: str) -> bool: """Verify file checksum.""" - sha256_hash = hashlib.sha256() - with open(file_path, "rb") as f: - for chunk in iter(lambda: f.read(8192), b""): - sha256_hash.update(chunk) - return sha256_hash.hexdigest() == expected_sha256 + return compute_sha256_file(file_path) == expected_sha256 def compute_checksum(self, file_path: Path) -> str: """Compute file SHA256.""" - sha256_hash = hashlib.sha256() - with open(file_path, "rb") as f: - for chunk in iter(lambda: f.read(8192), b""): - sha256_hash.update(chunk) - return sha256_hash.hexdigest() + return compute_sha256_file(file_path) async def list_models(self, query: str, task: str | None = None) -> list[HFRepo]: """Search for models on HuggingFace. @@ -128,37 +107,33 @@ async def list_models(self, query: str, task: str | None = None) -> list[HFRepo] Returns: List of matching :class:`HFRepo` instances. """ - import huggingface_hub.constants as hc from huggingface_hub import HfApi - was_offline = hc.HF_HUB_OFFLINE - hc.HF_HUB_OFFLINE = False try: - api = HfApi(token=self.token) - kwargs: dict[str, Any] = {"search": query, "sort": "downloads"} - if task is None: - pipeline_tags: list[str] = [] - elif task == "feature-extraction": - pipeline_tags = ["feature-extraction", "sentence-similarity"] - else: - pipeline_tags = [task] - seen: set[str] = set() - results: list[HFRepo] = [] - if pipeline_tags: - for tag in pipeline_tags: - kwargs["pipeline_tag"] = tag + with force_hf_online(): + api = HfApi(token=self.token) + kwargs: dict[str, Any] = {"search": query, "sort": "downloads"} + if task is None: + pipeline_tags: list[str] = [] + elif task == "feature-extraction": + pipeline_tags = ["feature-extraction", "sentence-similarity"] + else: + pipeline_tags = [task] + seen: set[str] = set() + results: list[HFRepo] = [] + if pipeline_tags: + for tag in pipeline_tags: + kwargs["pipeline_tag"] = tag + for r in api.list_models(**kwargs): + if r.id not in seen: + seen.add(r.id) + results.append(HFRepo.from_string(r.id)) + else: + kwargs.pop("pipeline_tag", None) for r in api.list_models(**kwargs): if r.id not in seen: seen.add(r.id) results.append(HFRepo.from_string(r.id)) - else: - kwargs.pop("pipeline_tag", None) - for r in api.list_models(**kwargs): - if r.id not in seen: - seen.add(r.id) - results.append(HFRepo.from_string(r.id)) - return results + return results except Exception as e: raise DownloadError(f"Failed to search models: {e}") from e - finally: - hc.HF_HUB_OFFLINE = was_offline diff --git a/src/application/models/embedding.py b/src/application/models/embedding.py index 315ade43..248ebc67 100644 --- a/src/application/models/embedding.py +++ b/src/application/models/embedding.py @@ -2,6 +2,7 @@ from __future__ import annotations +import asyncio import shutil from pathlib import Path from typing import Any @@ -118,7 +119,9 @@ async def download_model( raise DownloadError(f"Invalid repo_id '{repo_id}': {e}") from e temp_dir = self.embeddings_dir / "temp" / repo.owner / repo.name temp_dir.mkdir(parents=True, exist_ok=True) - downloaded_path = self.download_service.download_repo(repo, temp_dir) + downloaded_path = await asyncio.to_thread( + self.download_service.download_repo, repo, temp_dir + ) temp_path = Path(downloaded_path) final_dir = self.embeddings_dir / repo.owner / repo.name diff --git a/src/application/system/cache.py b/src/application/system/cache.py index 7eefc45c..9940e9f5 100644 --- a/src/application/system/cache.py +++ b/src/application/system/cache.py @@ -2,11 +2,12 @@ from __future__ import annotations -import hashlib import logging import time from typing import Any +from src.shared.utils.crypto_utils import compute_sha256_str + logger = logging.getLogger(__name__) DEFAULT_TTL = 300 # 5 minutes @@ -109,4 +110,4 @@ def generate_key(query: str, context: str = "") -> str: SHA-256 hash as cache key """ content = f"{query}:{context}" - return hashlib.sha256(content.encode()).hexdigest()[:16] + return compute_sha256_str(content)[:16] diff --git a/src/application/system/prompts.py b/src/application/system/prompts.py index d73f85b3..ad703b9d 100644 --- a/src/application/system/prompts.py +++ b/src/application/system/prompts.py @@ -50,7 +50,8 @@ def _save_all(prompts: dict[str, Prompt]) -> None: def _ensure_loaded() -> None: global _prompts - _prompts = _load_all() + if not _prompts: + _prompts = _load_all() def validate_name(name: str) -> None: diff --git a/src/application/tools/executor.py b/src/application/tools/executor.py index ec225899..c1cb0c5d 100644 --- a/src/application/tools/executor.py +++ b/src/application/tools/executor.py @@ -5,7 +5,7 @@ import logging from typing import Any -from .models import ToolDef, ToolExecutor, ToolResult +from .models import ToolContext, ToolDef, ToolExecutor, ToolResult logger = logging.getLogger(__name__) @@ -34,6 +34,7 @@ async def execute( tool_name: str, arguments: dict[str, Any], tool_call_id: str, + ctx: ToolContext | None = None, ) -> ToolResult: tool = self._tools.get(tool_name) if not tool: @@ -42,6 +43,9 @@ async def execute( f"Unknown tool '{tool_name}'. Available tools: {list(self._tools.keys())}", ) + arguments = {k: v for k, v in arguments.items() if k != "ctx"} + if ctx is not None and tool.has_ctx: + arguments["ctx"] = ctx try: result = await tool.fn(**arguments) return ToolResult(content=str(result), tool_call_id=tool_call_id) diff --git a/src/application/tools/models.py b/src/application/tools/models.py index c32fb437..02383af6 100644 --- a/src/application/tools/models.py +++ b/src/application/tools/models.py @@ -19,6 +19,7 @@ class ToolDef: description: str parameters: dict[str, Any] fn: Callable[..., Any] + has_ctx: bool = False def to_json_schema(self) -> dict[str, Any]: """Return the OpenAI-compatible tools JSON.""" @@ -48,6 +49,7 @@ async def execute( tool_name: str, arguments: dict[str, Any], tool_call_id: str, + ctx: ToolContext | None = None, ) -> ToolResult: raise NotImplementedError diff --git a/src/application/tools/registry.py b/src/application/tools/registry.py index 046e32c0..cfb2f251 100644 --- a/src/application/tools/registry.py +++ b/src/application/tools/registry.py @@ -50,8 +50,10 @@ def _build_json_schema(fn: Any) -> dict[str, Any]: sig = inspect.signature(fn) schema: dict[str, Any] = {"type": "object", "properties": {}, "required": []} + _skipped = {"self", "ctx"} + for name, param in sig.parameters.items(): - if name == "self": + if name in _skipped: continue annotation = param.annotation @@ -164,11 +166,13 @@ def decorator(fn: Any) -> ToolDef: first_sentence = doc.strip().split("\n")[0].strip() params["description"] = first_sentence + has_ctx = "ctx" in inspect.signature(fn).parameters td = ToolDef( name=fn.__name__, description=params.get("description", ""), parameters=params, fn=fn, + has_ctx=has_ctx, ) _tools.append(td) return td diff --git a/src/core/search/engine.py b/src/core/search/engine.py index 4dbb47d8..b4778329 100644 --- a/src/core/search/engine.py +++ b/src/core/search/engine.py @@ -74,10 +74,13 @@ def build_qdrant_filter(filters: dict[str, str]) -> Filter | None: def reset_search_embed_model() -> None: - """Reset the cached search embedding model singleton.""" + """Reset the cached search embedding model singleton and retriever cache.""" global _async_embed_model with _search_embed_model_lock: _async_embed_model = None + from src.core.search.retrievers import reset_retriever_cache + + reset_retriever_cache() def _get_search_embed_model() -> Any: diff --git a/src/core/search/retrievers.py b/src/core/search/retrievers.py index 11e379fe..73b9ef3a 100644 --- a/src/core/search/retrievers.py +++ b/src/core/search/retrievers.py @@ -3,6 +3,7 @@ from __future__ import annotations import logging +import threading from typing import Any from llama_index.core.retrievers import VectorIndexRetriever @@ -17,6 +18,33 @@ logger = logging.getLogger(__name__) +_async_client: AsyncQdrantClient | None = None +_async_client_lock = threading.Lock() + +_index_cache: dict[tuple[str, float], VectorStoreIndex] = {} +_index_cache_lock = threading.Lock() + + +def _get_async_client() -> AsyncQdrantClient: + global _async_client + if _async_client is not None: + return _async_client + with _async_client_lock: + if _async_client is not None: + return _async_client + cfg = get_config() + _async_client = AsyncQdrantClient(host=cfg.qdrant_host, port=cfg.qdrant_port) + return _async_client + + +def reset_retriever_cache() -> None: + """Clear cached indexes and the shared async client (call on model reset).""" + global _async_client + with _index_cache_lock: + _index_cache.clear() + with _async_client_lock: + _async_client = None + def _build_llama_filters( qdrant_filter: Any | None, @@ -44,6 +72,38 @@ def _build_llama_filters( return MetadataFilters(filters=filters_list) if filters_list else None # type: ignore[arg-type] +def _build_index(collection_name: str, alpha: float) -> VectorStoreIndex: + """Build (or retrieve from cache) a VectorStoreIndex for a collection.""" + cache_key = (collection_name, alpha) + with _index_cache_lock: + cached = _index_cache.get(cache_key) + if cached is not None: + return cached + + from src.core.search.engine import _get_search_embed_model + from src.core.search.sparse_encoder import create_sparse_encoder + + client = get_qdrant_client() + aclient = _get_async_client() + sparse_encoder = create_sparse_encoder() + vector_store = QdrantVectorStore( + client=client, + aclient=aclient, + collection_name=collection_name, + enable_hybrid=True, + sparse_doc_fn=sparse_encoder, + sparse_query_fn=sparse_encoder, + sparse_vector_name="text-sparse", + ) + + embed_model = _get_search_embed_model() + index = VectorStoreIndex.from_vector_store(vector_store, embed_model=embed_model) + + with _index_cache_lock: + _index_cache[cache_key] = index + return index + + def get_collection_retriever( collection_name: str, top_k: int = 30, @@ -52,6 +112,10 @@ def get_collection_retriever( ) -> VectorIndexRetriever: """Get a LlamaIndex retriever for a specific Qdrant collection. + The underlying VectorStoreIndex is cached per (collection, alpha) pair. + A lightweight VectorIndexRetriever is created per call with the + per-query top_k and metadata_filter. + Args: collection_name: Qdrant collection name. top_k: Number of results per collection (before fusion). @@ -61,28 +125,8 @@ def get_collection_retriever( Returns: Configured VectorIndexRetriever. """ - from src.core.search.engine import _get_search_embed_model - try: - from src.core.search.sparse_encoder import create_sparse_encoder - - client = get_qdrant_client() - cfg = get_config() - aclient = AsyncQdrantClient(host=cfg.qdrant_host, port=cfg.qdrant_port) - sparse_encoder = create_sparse_encoder() - vector_store = QdrantVectorStore( - client=client, - aclient=aclient, - collection_name=collection_name, - enable_hybrid=True, - sparse_doc_fn=sparse_encoder, - sparse_query_fn=sparse_encoder, - sparse_vector_name="text-sparse", - ) - - embed_model = _get_search_embed_model() - index = VectorStoreIndex.from_vector_store(vector_store, embed_model=embed_model) - + index = _build_index(collection_name, alpha) llama_filters = _build_llama_filters(metadata_filter) retriever = VectorIndexRetriever( diff --git a/src/infrastructure/database/core.py b/src/infrastructure/database/core.py index c20cff0e..c6eaf809 100644 --- a/src/infrastructure/database/core.py +++ b/src/infrastructure/database/core.py @@ -208,7 +208,11 @@ def get_config(self, key: str) -> Any | None: return None def set_config( - self, key: str, value: dict[str, Any] | list[Any] | str | int | bool | None + self, + key: str, + value: dict[str, Any] | list[Any] | str | int | bool | None, + *, + persist: bool = True, ) -> None: with self._lock: self._writer_conn.execute( @@ -216,6 +220,7 @@ def set_config( (key, json.dumps(value)), ) self._writer_conn.commit() + if persist: try: self.persist() except Exception: diff --git a/src/infrastructure/database/doc_ops.py b/src/infrastructure/database/doc_ops.py index 7b1b1368..22b31ae0 100644 --- a/src/infrastructure/database/doc_ops.py +++ b/src/infrastructure/database/doc_ops.py @@ -2,13 +2,13 @@ from __future__ import annotations -import hashlib import logging import os from datetime import datetime, timezone, timedelta from pathlib import Path from src.shared.constants import NULL_UUID +from src.shared.utils.crypto_utils import compute_sha256_file from typing import TYPE_CHECKING, Any if TYPE_CHECKING: @@ -504,18 +504,11 @@ def resync_local_file_sizes(self, base_path: str) -> dict[str, int]: return {"updated": 0, "skipped": 0, "error": 0, "incomplete": 0} def _hash_file(path: Path) -> str | None: - h = hashlib.sha256() - try: - with open(path, "rb") as f: - while True: - chunk = f.read(8192) - if not chunk: - break - h.update(chunk) - return h.hexdigest() - except OSError as e: - logger.warning(f"[resync_local_file_sizes] Cannot read {path}: {e}") + digest = compute_sha256_file(path) + if not digest: + logger.warning(f"[resync_local_file_sizes] Cannot read {path}") return None + return digest snapshot: list[tuple[str, str | None, str | None]] = [] with self._lock: diff --git a/src/infrastructure/github/api.py b/src/infrastructure/github/api.py index 8ef7bf41..9c2c3974 100644 --- a/src/infrastructure/github/api.py +++ b/src/infrastructure/github/api.py @@ -2,46 +2,45 @@ from __future__ import annotations -import httpx +from src.shared.http import get_async_pooled_client GITHUB_API_URL = "https://api.github.com/repos" -async def list_releases(owner: str, repo: str, github_token: str | None = None) -> list[dict]: - """List all releases for a repository.""" - url = f"{GITHUB_API_URL}/{owner}/{repo}/releases" - headers = {"Accept": "application/vnd.github+json"} +def _github_headers(github_token: str | None = None) -> dict[str, str]: + headers: dict[str, str] = {"Accept": "application/vnd.github+json"} if github_token: headers["Authorization"] = f"Bearer {github_token}" + return headers - async with httpx.AsyncClient(timeout=30.0) as client: - response = await client.get(url, headers=headers) - response.raise_for_status() - data = response.json() - return [ - { - "tag_name": r.get("tag_name"), - "name": r.get("name"), - "published_at": r.get("published_at"), - "prerelease": r.get("prerelease"), - "draft": r.get("draft"), - "assets_count": len(r.get("assets", [])), - } - for r in data - ] + +async def list_releases(owner: str, repo: str, github_token: str | None = None) -> list[dict]: + """List all releases for a repository.""" + url = f"{GITHUB_API_URL}/{owner}/{repo}/releases" + client = get_async_pooled_client(timeout=30.0) + response = await client.get(url, headers=_github_headers(github_token)) + response.raise_for_status() + data = response.json() + return [ + { + "tag_name": r.get("tag_name"), + "name": r.get("name"), + "published_at": r.get("published_at"), + "prerelease": r.get("prerelease"), + "draft": r.get("draft"), + "assets_count": len(r.get("assets", [])), + } + for r in data + ] async def info_release(owner: str, repo: str, tag: str, github_token: str | None = None) -> dict: """Get release info by tag.""" url = f"{GITHUB_API_URL}/{owner}/{repo}/releases/tags/{tag}" - headers = {"Accept": "application/vnd.github+json"} - if github_token: - headers["Authorization"] = f"Bearer {github_token}" - - async with httpx.AsyncClient(timeout=30.0) as client: - response = await client.get(url, headers=headers) - response.raise_for_status() - return response.json() # type: ignore[no-any-return] + client = get_async_pooled_client(timeout=30.0) + response = await client.get(url, headers=_github_headers(github_token)) + response.raise_for_status() + return response.json() # type: ignore[no-any-return] async def list_release_files( @@ -49,23 +48,19 @@ async def list_release_files( ) -> list[dict]: """List all files (assets) of a release.""" url = f"{GITHUB_API_URL}/{owner}/{repo}/releases/tags/{tag}" - headers = {"Accept": "application/vnd.github+json"} - if github_token: - headers["Authorization"] = f"Bearer {github_token}" - - async with httpx.AsyncClient(timeout=30.0) as client: - response = await client.get(url, headers=headers) - response.raise_for_status() - data = response.json() - return [ - { - "name": a["name"], - "size": a["size"], - "download_url": a["browser_download_url"], - "content_type": a["content_type"], - } - for a in data.get("assets", []) - ] + client = get_async_pooled_client(timeout=30.0) + response = await client.get(url, headers=_github_headers(github_token)) + response.raise_for_status() + data = response.json() + return [ + { + "name": a["name"], + "size": a["size"], + "download_url": a["browser_download_url"], + "content_type": a["content_type"], + } + for a in data.get("assets", []) + ] async def download_release_file( @@ -77,22 +72,18 @@ async def download_release_file( ) -> dict: """Download a specific file from a release.""" url = f"{GITHUB_API_URL}/{owner}/{repo}/releases/tags/{tag}" - headers = {"Accept": "application/vnd.github+json"} - if github_token: - headers["Authorization"] = f"Bearer {github_token}" - - async with httpx.AsyncClient(timeout=30.0) as client: - response = await client.get(url, headers=headers) - response.raise_for_status() - data = response.json() - - for asset in data.get("assets", []): - if asset["name"] == filename: - return { - "name": asset["name"], - "size": asset["size"], - "download_url": asset["browser_download_url"], - "content_type": asset["content_type"], - } + client = get_async_pooled_client(timeout=30.0) + response = await client.get(url, headers=_github_headers(github_token)) + response.raise_for_status() + data = response.json() + + for asset in data.get("assets", []): + if asset["name"] == filename: + return { + "name": asset["name"], + "size": asset["size"], + "download_url": asset["browser_download_url"], + "content_type": asset["content_type"], + } - raise ValueError(f"File '{filename}' not found in release '{tag}'") + raise ValueError(f"File '{filename}' not found in release '{tag}'") diff --git a/src/infrastructure/github/git.py b/src/infrastructure/github/git.py index 524cd005..c1dbe29e 100644 --- a/src/infrastructure/github/git.py +++ b/src/infrastructure/github/git.py @@ -67,10 +67,10 @@ def _validate_git_url(url: str) -> None: raise ValueError("URL points to localhost, which is not allowed") try: ip = ipaddress.ip_address(host) - if ip.is_private or ip.is_loopback or ip.is_link_local: - raise ValueError(f"URL points to a private/reserved IP address: {host}") except ValueError: - pass + return + if ip.is_private or ip.is_loopback or ip.is_link_local: + raise ValueError(f"URL points to a private/reserved IP address: {host}") def clone_repo( @@ -272,13 +272,18 @@ def delete_repo(org: str, name: str, repos_dir: Path | None = None) -> dict[str, def list_repos( - repos_dir: Path | None = None, org_filter: str | None = None + repos_dir: Path | None = None, + org_filter: str | None = None, + *, + fetch_remote: bool = False, ) -> list[dict[str, Any]]: """List all cloned repositories with their metadata. Args: repos_dir: Base directory for cloned repos. org_filter: If provided, only list repos under this org directory. + fetch_remote: If True, perform a network fetch per repo to get + the latest remote HEAD. Default False (local refs only). """ repos_dir = Path(repos_dir or get_config().paths_github_dir).resolve() repos_dir.mkdir(parents=True, exist_ok=True) @@ -297,23 +302,22 @@ def list_repos( "name": repo_dir.name, "path": str(repo_dir), } - # gitpython provides branches, active_branch, remote().url info["branch"] = repo.active_branch.name if repo.active_branch else None try: info["remote_url"] = repo.remote().url except Exception: info["remote_url"] = None - try: - origin = repo.remotes.origin - origin.fetch() - remote_ref = f"origin/{info.get('branch')}" - remote_head = ( - repo.refs[remote_ref].commit.hexsha if remote_ref in repo.refs else "" - ) - info["remote_head"] = remote_head - except Exception as e: - logger.warning("Fetch failed for %s/%s: %s", org_dir.name, repo_dir.name, e) - info["remote_head"] = "" + info["remote_head"] = "" + if fetch_remote: + try: + origin = repo.remotes.origin + origin.fetch() + remote_ref = f"origin/{info.get('branch')}" + info["remote_head"] = ( + repo.refs[remote_ref].commit.hexsha if remote_ref in repo.refs else "" + ) + except Exception as e: + logger.warning("Fetch failed for %s/%s: %s", org_dir.name, repo_dir.name, e) repos.append(info) return sorted(repos, key=lambda r: (r["org"], r["name"])) diff --git a/src/main.py b/src/main.py index c71da14a..a6ee7d12 100644 --- a/src/main.py +++ b/src/main.py @@ -61,6 +61,7 @@ from src.api.v1.documents.spec import router as spec_v1_router from src.api.v1.system.system_prompt import router as prompts_v1_router from src.api.v1.sigma.translate import router as translate_v1_router +from src.shared.http import close_all_async_pooled_clients from src.shared.service_manager import shutdown_all_services from src.config.settings import TEMP_DIR from src.presentation import STATIC_DIR @@ -232,7 +233,7 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: logger.info("Config initialized.") # Start the background task dispatcher in its own thread - dispatcher = TaskDispatcher(poll_interval=1, max_workers=4) + dispatcher = TaskDispatcher(poll_interval=0.2, max_workers=4) app.state.dispatcher = dispatcher dispatcher.start() logger.info("Dispatcher started in background thread.") @@ -273,6 +274,7 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: if dispatcher: dispatcher.stop() logger.info("Dispatcher stopped.") + await close_all_async_pooled_clients() await stop_llamacpp() await shutdown_all_services() await stop_qdrant() diff --git a/src/shared/http.py b/src/shared/http.py index a136ff20..8e907330 100644 --- a/src/shared/http.py +++ b/src/shared/http.py @@ -90,6 +90,60 @@ def close_all_pooled_clients() -> None: _pool.clear() +# ------------------------------------------------------------------ +# Async connection pool +# ------------------------------------------------------------------ + +_async_pool: dict[str, httpx.AsyncClient] = {} + + +def get_async_pooled_client( + timeout: float = DEFAULT_TIMEOUT, + headers: dict[str, str] | None = None, + follow_redirects: bool = True, +) -> httpx.AsyncClient: + """Return a pooled ``httpx.AsyncClient`` keyed by (timeout, follow_redirects). + + Must be called from within a running event loop. Connections are reused + across requests, reducing TCP handshake overhead for repeated API calls. + """ + key = _pool_key(timeout, follow_redirects) + client = _async_pool.get(key) + if client is not None: + return client + + merged: dict[str, str] = {"User-Agent": DEFAULT_USER_AGENT} + if headers: + merged.update(headers) + + transport = httpx.AsyncHTTPTransport( + limits=httpx.Limits( + max_connections=100, + max_keepalive_connections=20, + keepalive_expiry=30.0, + ), + ) + + new_client = httpx.AsyncClient( + timeout=httpx.Timeout(timeout), + headers=merged, + follow_redirects=follow_redirects, + transport=transport, + ) + _async_pool[key] = new_client + return new_client + + +async def close_all_async_pooled_clients() -> None: + """Close all pooled async HTTP clients. Call at shutdown.""" + for client in _async_pool.values(): + try: + await client.aclose() + except Exception: + pass + _async_pool.clear() + + def create_client( timeout: float = DEFAULT_TIMEOUT, headers: dict[str, str] | None = None, @@ -197,10 +251,12 @@ def download_file( for attempt in range(1, max_retries + 1): try: - resp = client.get(url) - resp.raise_for_status() path.parent.mkdir(parents=True, exist_ok=True) - path.write_bytes(resp.content) + with client.stream("GET", url) as resp: + resp.raise_for_status() + with open(path, "wb") as f: + for chunk in resp.iter_bytes(chunk_size=1024 * 1024): + f.write(chunk) return True, None except OSError as exc: diff --git a/src/shared/session.py b/src/shared/session.py index 7e66d5e0..0855d6fa 100644 --- a/src/shared/session.py +++ b/src/shared/session.py @@ -41,5 +41,5 @@ def _touch(self, session_id: str) -> None: self._accessed[session_id] = time.monotonic() def _evict(self) -> None: - oldest = min(self._accessed, key=self._accessed.get) + oldest = min(self._accessed, key=lambda k: self._accessed[k]) self.delete(oldest) diff --git a/src/shared/utils/hf_hub.py b/src/shared/utils/hf_hub.py new file mode 100644 index 00000000..565e9798 --- /dev/null +++ b/src/shared/utils/hf_hub.py @@ -0,0 +1,24 @@ +"""Context manager to temporarily force HuggingFace Hub online mode.""" + +from __future__ import annotations + +from contextlib import contextmanager +from typing import Iterator + + +@contextmanager +def force_hf_online() -> Iterator[None]: + """Temporarily disable HF_HUB_OFFLINE for network operations. + + Usage: + with force_hf_online(): + snapshot_download(...) + """ + import huggingface_hub.constants as hc + + was_offline = hc.HF_HUB_OFFLINE + hc.HF_HUB_OFFLINE = False + try: + yield + finally: + hc.HF_HUB_OFFLINE = was_offline diff --git a/src/workers/processor.py b/src/workers/processor.py index 6f5d7668..bbaef5c0 100644 --- a/src/workers/processor.py +++ b/src/workers/processor.py @@ -31,7 +31,7 @@ class TaskDispatcher: WorkerName.MODEL_SYNC: ModelSyncWorker, } - def __init__(self, poll_interval: float = 1.0, max_workers: int = 1): + def __init__(self, poll_interval: float = 0.2, max_workers: int = 4): self.poll_interval = poll_interval self.max_workers = max_workers self._running = False diff --git a/tests/unit/application/tools/test_ctx_injection.py b/tests/unit/application/tools/test_ctx_injection.py new file mode 100644 index 00000000..3895babf --- /dev/null +++ b/tests/unit/application/tools/test_ctx_injection.py @@ -0,0 +1,131 @@ +"""Tests for ToolContext injection into tool execution.""" + +import pytest + +from src.application.tools.models import ToolContext +from src.application.tools.executor import ToolDispatcher +from src.application.tools.registry import reset_tools, tool + + +class MockSearchEngine: + async def search(self, query: str, **kwargs) -> list[dict]: + return [{"text": "mock result", "score": 0.95, "metadata": {"title": "Test Rule"}}] + + +class MockLLMClient: + async def chat(self, messages, **kwargs) -> str: + return "mock llm response" + + +class TestCtxInjection: + """Verify ctx is injected into tools and excluded from JSON schema.""" + + def setup_method(self) -> None: + reset_tools() + + def test_ctx_excluded_from_json_schema(self) -> None: + """ctx must not appear in the tool schema sent to the LLM.""" + + @tool + async def test_tool(query: str, ctx: ToolContext | None = None) -> str: + """A tool with ctx. + + :param query: The query. + """ + return "ok" + + schema = test_tool.parameters + assert "ctx" not in schema["properties"] + assert "ctx" not in schema["required"] + assert "query" in schema["properties"] + + def test_self_excluded_from_json_schema(self) -> None: + """self must not appear in the tool schema.""" + + class MyTools: + @tool + async def bound_tool(self, query: str) -> str: + """A bound method tool. + + :param query: The query. + """ + return "ok" + + schema = MyTools.bound_tool.parameters + assert "self" not in schema["properties"] + + @pytest.mark.asyncio + async def test_ctx_injected_into_tool(self) -> None: + """ctx is passed to the tool function when provided to execute().""" + received_ctx: ToolContext | None = None + + @tool + async def ctx_tool(query: str, ctx: ToolContext | None = None) -> str: + """A tool that captures ctx. + + :param query: The query. + """ + nonlocal received_ctx + received_ctx = ctx + return "ctx received" if ctx else "no ctx" + + ctx = ToolContext(search_engine=MockSearchEngine(), llm_client=MockLLMClient()) + dispatcher = ToolDispatcher([ctx_tool]) + result = await dispatcher.execute("ctx_tool", {"query": "test"}, "call_1", ctx=ctx) + + assert result.content == "ctx received" + assert received_ctx is ctx + + @pytest.mark.asyncio + async def test_ctx_none_when_not_provided(self) -> None: + """ctx is None when execute() is called without ctx.""" + + @tool + async def no_ctx_tool(query: str, ctx: ToolContext | None = None) -> str: + """A tool checking ctx is None. + + :param query: The query. + """ + return "has ctx" if ctx else "no ctx" + + dispatcher = ToolDispatcher([no_ctx_tool]) + result = await dispatcher.execute("no_ctx_tool", {"query": "test"}, "call_1") + assert result.content == "no ctx" + + @pytest.mark.asyncio + async def test_ctx_stripped_from_arguments(self) -> None: + """Even if LLM sends ctx in arguments, it is stripped before call.""" + received_ctx: ToolContext | None = None + + @tool + async def strip_tool(query: str, ctx: ToolContext | None = None) -> str: + """A tool. + + :param query: The query. + """ + nonlocal received_ctx + received_ctx = ctx + return "ok" + + real_ctx = ToolContext(search_engine=MockSearchEngine(), llm_client=MockLLMClient()) + dispatcher = ToolDispatcher([strip_tool]) + result = await dispatcher.execute( + "strip_tool", {"query": "test", "ctx": "bogus_llm_value"}, "call_1", ctx=real_ctx + ) + assert result.content == "ok" + assert received_ctx is real_ctx + + @pytest.mark.asyncio + async def test_sigma_tools_receive_ctx(self) -> None: + """All 5 registered sigma tools work when ctx is injected.""" + from src.application.tools.sigma.tools import search_sigma # noqa: F401 + + ctx = ToolContext(search_engine=MockSearchEngine(), llm_client=MockLLMClient()) + dispatcher = ToolDispatcher() + dispatcher.register(search_sigma) + + result = await dispatcher.execute( + "search_sigma", {"query": "suspicious process"}, "tc1", ctx=ctx + ) + assert "Error: tool context not available" not in result.content + assert "Test Rule" in result.content diff --git a/tests/unit/back/github/test_api.py b/tests/unit/back/github/test_api.py index 1c50486c..2bfd5062 100644 --- a/tests/unit/back/github/test_api.py +++ b/tests/unit/back/github/test_api.py @@ -1,228 +1,229 @@ -"""Tests for GitHub API interactions.""" - -from unittest.mock import AsyncMock, MagicMock, patch - -import httpx -import pytest - -from src.infrastructure.github.api import ( - download_release_file, - info_release, - list_release_files, - list_releases, -) - - -def _mock_response(data): - """Create a mock httpx response with synchronous json() method.""" - resp = MagicMock() - resp.json.return_value = data - resp.raise_for_status.return_value = None - return resp - - -@pytest.fixture -def mock_client(): - client = AsyncMock() - client.__aenter__.return_value = client - return client - - -@pytest.fixture -def sample_releases() -> list[dict]: - return [ - { - "tag_name": "v1.0.0", - "name": "Release 1.0.0", - "published_at": "2024-01-01T00:00:00Z", - "prerelease": False, - "draft": False, - "assets": [ - {"name": "asset1.zip", "size": 100, "browser_download_url": "https://example.com/1"} - ], - } - ] - - -@pytest.fixture -def sample_release_detail() -> dict: - return { - "tag_name": "v1.0.0", - "name": "Release 1.0.0", - "assets": [ - { - "name": "asset1.zip", - "size": 100, - "browser_download_url": "https://example.com/1", - "content_type": "application/zip", - } - ], - } - - -class TestListReleases: - @pytest.mark.asyncio - async def test_success(self, mock_client: AsyncMock, sample_releases: list[dict]) -> None: - mock_client.get.return_value = _mock_response(sample_releases) - with patch("httpx.AsyncClient", return_value=mock_client): - result = await list_releases("owner", "repo") - assert len(result) == 1 - assert result[0]["tag_name"] == "v1.0.0" - assert result[0]["assets_count"] == 1 - - @pytest.mark.asyncio - async def test_with_token(self, mock_client: AsyncMock, sample_releases: list[dict]) -> None: - mock_client.get.return_value = _mock_response(sample_releases) - with patch("httpx.AsyncClient", return_value=mock_client): - await list_releases("owner", "repo", github_token="token-123") - call_kwargs = mock_client.get.call_args - headers = call_kwargs[1]["headers"] - assert "Authorization" in headers - assert headers["Authorization"] == "Bearer token-123" - - @pytest.mark.asyncio - async def test_raises_on_http_error(self, mock_client: AsyncMock) -> None: - mock_response = MagicMock() - mock_response.raise_for_status.side_effect = httpx.HTTPStatusError( - "404", request=MagicMock(), response=MagicMock() - ) - mock_client.get.return_value = mock_response - with patch("httpx.AsyncClient", return_value=mock_client): - with pytest.raises(httpx.HTTPStatusError): - await list_releases("owner", "repo") - - -class TestInfoRelease: - @pytest.mark.asyncio - async def test_success(self, mock_client: AsyncMock, sample_release_detail: dict) -> None: - mock_client.get.return_value = _mock_response(sample_release_detail) - with patch("httpx.AsyncClient", return_value=mock_client): - result = await info_release("owner", "repo", "v1.0.0") - assert result["tag_name"] == "v1.0.0" - - @pytest.mark.asyncio - async def test_uses_tag_as_is_without_v_prefix( - self, mock_client: AsyncMock, sample_release_detail: dict - ) -> None: - mock_client.get.return_value = _mock_response(sample_release_detail) - with patch("httpx.AsyncClient", return_value=mock_client): - await info_release("owner", "repo", "1.0.0") - call_url = mock_client.get.call_args[0][0] - assert "1.0.0" in call_url - assert "v1.0.0" not in call_url - - @pytest.mark.asyncio - async def test_handles_b_prefix_tag( - self, mock_client: AsyncMock, sample_release_detail: dict - ) -> None: - mock_client.get.return_value = _mock_response(sample_release_detail) - with patch("httpx.AsyncClient", return_value=mock_client): - await info_release("owner", "repo", "b9601") - call_url = mock_client.get.call_args[0][0] - assert "b9601" in call_url - - @pytest.mark.asyncio - async def test_preserves_v_prefix( - self, mock_client: AsyncMock, sample_release_detail: dict - ) -> None: - mock_client.get.return_value = _mock_response(sample_release_detail) - with patch("httpx.AsyncClient", return_value=mock_client): - await info_release("owner", "repo", "v2.0.0") - call_url = mock_client.get.call_args[0][0] - assert "v2.0.0" in call_url - - @pytest.mark.asyncio - async def test_with_token(self, mock_client: AsyncMock, sample_release_detail: dict) -> None: - mock_client.get.return_value = _mock_response(sample_release_detail) - with patch("httpx.AsyncClient", return_value=mock_client): - await info_release("owner", "repo", "v1.0.0", github_token="token-789") - call_kwargs = mock_client.get.call_args - headers = call_kwargs[1]["headers"] - assert headers["Authorization"] == "Bearer token-789" - - -class TestListReleaseFiles: - @pytest.mark.asyncio - async def test_success(self, mock_client: AsyncMock) -> None: - data = { - "tag_name": "v1.0.0", - "assets": [ - { - "name": "file1.zip", - "size": 100, - "browser_download_url": "https://example.com/1", - "content_type": "application/zip", - }, - { - "name": "file2.tar.gz", - "size": 200, - "browser_download_url": "https://example.com/2", - "content_type": "application/gzip", - }, - ], - } - mock_client.get.return_value = _mock_response(data) - with patch("httpx.AsyncClient", return_value=mock_client): - result = await list_release_files("owner", "repo", "v1.0.0") - assert len(result) == 2 - assert result[0]["name"] == "file1.zip" - assert result[1]["name"] == "file2.tar.gz" - - @pytest.mark.asyncio - async def test_returns_empty_when_no_assets(self, mock_client: AsyncMock) -> None: - data = {"tag_name": "v1.0.0", "assets": []} - mock_client.get.return_value = _mock_response(data) - with patch("httpx.AsyncClient", return_value=mock_client): - result = await list_release_files("owner", "repo", "v1.0.0") - assert result == [] - - @pytest.mark.asyncio - async def test_with_token(self, mock_client: AsyncMock) -> None: - data = {"tag_name": "v1.0.0", "assets": []} - mock_client.get.return_value = _mock_response(data) - with patch("httpx.AsyncClient", return_value=mock_client): - await list_release_files("owner", "repo", "v1.0.0", github_token="token-abc") - call_kwargs = mock_client.get.call_args - headers = call_kwargs[1]["headers"] - assert headers["Authorization"] == "Bearer token-abc" - - -class TestDownloadReleaseFile: - @pytest.mark.asyncio - async def test_finds_matching_asset(self, mock_client: AsyncMock) -> None: - data = { - "tag_name": "v1.0.0", - "assets": [ - { - "name": "target.zip", - "size": 500, - "browser_download_url": "https://example.com/target", - "content_type": "application/zip", - } - ], - } - mock_client.get.return_value = _mock_response(data) - with patch("httpx.AsyncClient", return_value=mock_client): - result = await download_release_file("owner", "repo", "target.zip", "v1.0.0") - assert result["name"] == "target.zip" - assert result["size"] == 500 - - @pytest.mark.asyncio - async def test_raises_when_not_found(self, mock_client: AsyncMock) -> None: - data = {"tag_name": "v1.0.0", "assets": []} - mock_client.get.return_value = _mock_response(data) - with patch("httpx.AsyncClient", return_value=mock_client): - with pytest.raises(ValueError, match="not found"): - await download_release_file("owner", "repo", "missing.zip", "v1.0.0") - - @pytest.mark.asyncio - async def test_with_token(self, mock_client: AsyncMock, sample_release_detail: dict) -> None: - mock_client.get.return_value = _mock_response(sample_release_detail) - with patch("httpx.AsyncClient", return_value=mock_client): - result = await download_release_file( - "owner", "repo", "asset1.zip", "v1.0.0", github_token="token-456" - ) - assert result["name"] == "asset1.zip" - call_kwargs = mock_client.get.call_args - headers = call_kwargs[1]["headers"] - assert headers["Authorization"] == "Bearer token-456" +"""Tests for GitHub API interactions.""" + +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest + +from src.infrastructure.github.api import ( + download_release_file, + info_release, + list_release_files, + list_releases, +) + +MOCK_TARGET = "src.infrastructure.github.api.get_async_pooled_client" + + +def _mock_response(data): + """Create a mock httpx response with synchronous json() method.""" + resp = MagicMock() + resp.json.return_value = data + resp.raise_for_status.return_value = None + return resp + + +@pytest.fixture +def mock_client(): + client = AsyncMock() + return client + + +@pytest.fixture +def sample_releases() -> list[dict]: + return [ + { + "tag_name": "v1.0.0", + "name": "Release 1.0.0", + "published_at": "2024-01-01T00:00:00Z", + "prerelease": False, + "draft": False, + "assets": [ + {"name": "asset1.zip", "size": 100, "browser_download_url": "https://example.com/1"} + ], + } + ] + + +@pytest.fixture +def sample_release_detail() -> dict: + return { + "tag_name": "v1.0.0", + "name": "Release 1.0.0", + "assets": [ + { + "name": "asset1.zip", + "size": 100, + "browser_download_url": "https://example.com/1", + "content_type": "application/zip", + } + ], + } + + +class TestListReleases: + @pytest.mark.asyncio + async def test_success(self, mock_client: AsyncMock, sample_releases: list[dict]) -> None: + mock_client.get.return_value = _mock_response(sample_releases) + with patch(MOCK_TARGET, return_value=mock_client): + result = await list_releases("owner", "repo") + assert len(result) == 1 + assert result[0]["tag_name"] == "v1.0.0" + assert result[0]["assets_count"] == 1 + + @pytest.mark.asyncio + async def test_with_token(self, mock_client: AsyncMock, sample_releases: list[dict]) -> None: + mock_client.get.return_value = _mock_response(sample_releases) + with patch(MOCK_TARGET, return_value=mock_client): + await list_releases("owner", "repo", github_token="token-123") + call_kwargs = mock_client.get.call_args + headers = call_kwargs[1]["headers"] + assert "Authorization" in headers + assert headers["Authorization"] == "Bearer token-123" + + @pytest.mark.asyncio + async def test_raises_on_http_error(self, mock_client: AsyncMock) -> None: + mock_response = MagicMock() + mock_response.raise_for_status.side_effect = httpx.HTTPStatusError( + "404", request=MagicMock(), response=MagicMock() + ) + mock_client.get.return_value = mock_response + with patch(MOCK_TARGET, return_value=mock_client): + with pytest.raises(httpx.HTTPStatusError): + await list_releases("owner", "repo") + + +class TestInfoRelease: + @pytest.mark.asyncio + async def test_success(self, mock_client: AsyncMock, sample_release_detail: dict) -> None: + mock_client.get.return_value = _mock_response(sample_release_detail) + with patch(MOCK_TARGET, return_value=mock_client): + result = await info_release("owner", "repo", "v1.0.0") + assert result["tag_name"] == "v1.0.0" + + @pytest.mark.asyncio + async def test_uses_tag_as_is_without_v_prefix( + self, mock_client: AsyncMock, sample_release_detail: dict + ) -> None: + mock_client.get.return_value = _mock_response(sample_release_detail) + with patch(MOCK_TARGET, return_value=mock_client): + await info_release("owner", "repo", "1.0.0") + call_url = mock_client.get.call_args[0][0] + assert "1.0.0" in call_url + assert "v1.0.0" not in call_url + + @pytest.mark.asyncio + async def test_handles_b_prefix_tag( + self, mock_client: AsyncMock, sample_release_detail: dict + ) -> None: + mock_client.get.return_value = _mock_response(sample_release_detail) + with patch(MOCK_TARGET, return_value=mock_client): + await info_release("owner", "repo", "b9601") + call_url = mock_client.get.call_args[0][0] + assert "b9601" in call_url + + @pytest.mark.asyncio + async def test_preserves_v_prefix( + self, mock_client: AsyncMock, sample_release_detail: dict + ) -> None: + mock_client.get.return_value = _mock_response(sample_release_detail) + with patch(MOCK_TARGET, return_value=mock_client): + await info_release("owner", "repo", "v2.0.0") + call_url = mock_client.get.call_args[0][0] + assert "v2.0.0" in call_url + + @pytest.mark.asyncio + async def test_with_token(self, mock_client: AsyncMock, sample_release_detail: dict) -> None: + mock_client.get.return_value = _mock_response(sample_release_detail) + with patch(MOCK_TARGET, return_value=mock_client): + await info_release("owner", "repo", "v1.0.0", github_token="token-789") + call_kwargs = mock_client.get.call_args + headers = call_kwargs[1]["headers"] + assert headers["Authorization"] == "Bearer token-789" + + +class TestListReleaseFiles: + @pytest.mark.asyncio + async def test_success(self, mock_client: AsyncMock) -> None: + data = { + "tag_name": "v1.0.0", + "assets": [ + { + "name": "file1.zip", + "size": 100, + "browser_download_url": "https://example.com/1", + "content_type": "application/zip", + }, + { + "name": "file2.tar.gz", + "size": 200, + "browser_download_url": "https://example.com/2", + "content_type": "application/gzip", + }, + ], + } + mock_client.get.return_value = _mock_response(data) + with patch(MOCK_TARGET, return_value=mock_client): + result = await list_release_files("owner", "repo", "v1.0.0") + assert len(result) == 2 + assert result[0]["name"] == "file1.zip" + assert result[1]["name"] == "file2.tar.gz" + + @pytest.mark.asyncio + async def test_returns_empty_when_no_assets(self, mock_client: AsyncMock) -> None: + data = {"tag_name": "v1.0.0", "assets": []} + mock_client.get.return_value = _mock_response(data) + with patch(MOCK_TARGET, return_value=mock_client): + result = await list_release_files("owner", "repo", "v1.0.0") + assert result == [] + + @pytest.mark.asyncio + async def test_with_token(self, mock_client: AsyncMock) -> None: + data = {"tag_name": "v1.0.0", "assets": []} + mock_client.get.return_value = _mock_response(data) + with patch(MOCK_TARGET, return_value=mock_client): + await list_release_files("owner", "repo", "v1.0.0", github_token="token-abc") + call_kwargs = mock_client.get.call_args + headers = call_kwargs[1]["headers"] + assert headers["Authorization"] == "Bearer token-abc" + + +class TestDownloadReleaseFile: + @pytest.mark.asyncio + async def test_finds_matching_asset(self, mock_client: AsyncMock) -> None: + data = { + "tag_name": "v1.0.0", + "assets": [ + { + "name": "target.zip", + "size": 500, + "browser_download_url": "https://example.com/target", + "content_type": "application/zip", + } + ], + } + mock_client.get.return_value = _mock_response(data) + with patch(MOCK_TARGET, return_value=mock_client): + result = await download_release_file("owner", "repo", "target.zip", "v1.0.0") + assert result["name"] == "target.zip" + assert result["size"] == 500 + + @pytest.mark.asyncio + async def test_raises_when_not_found(self, mock_client: AsyncMock) -> None: + data = {"tag_name": "v1.0.0", "assets": []} + mock_client.get.return_value = _mock_response(data) + with patch(MOCK_TARGET, return_value=mock_client): + with pytest.raises(ValueError, match="not found"): + await download_release_file("owner", "repo", "missing.zip", "v1.0.0") + + @pytest.mark.asyncio + async def test_with_token(self, mock_client: AsyncMock, sample_release_detail: dict) -> None: + mock_client.get.return_value = _mock_response(sample_release_detail) + with patch(MOCK_TARGET, return_value=mock_client): + result = await download_release_file( + "owner", "repo", "asset1.zip", "v1.0.0", github_token="token-456" + ) + assert result["name"] == "asset1.zip" + call_kwargs = mock_client.get.call_args + headers = call_kwargs[1]["headers"] + assert headers["Authorization"] == "Bearer token-456" diff --git a/tests/unit/back/github/test_git.py b/tests/unit/back/github/test_git.py index a1cb2cc7..bf606b74 100644 --- a/tests/unit/back/github/test_git.py +++ b/tests/unit/back/github/test_git.py @@ -43,10 +43,16 @@ def test_local_ip_blocked(self) -> None: with pytest.raises(ValueError, match="localhost"): _validate_git_url("https://127.0.0.1/repo.git") - def test_private_ip_not_blocked_due_to_bug(self) -> None: - # NOTE: The except ValueError: pass in _validate_git_url catches the - # intentional raise for private/reserved IPs, so they are NOT blocked. - _validate_git_url("https://10.0.0.1/repo.git") + def test_private_ip_blocked(self) -> None: + with pytest.raises(ValueError, match="private/reserved IP"): + _validate_git_url("https://10.0.0.1/repo.git") + + def test_private_ip_192_blocked(self) -> None: + with pytest.raises(ValueError, match="private/reserved IP"): + _validate_git_url("https://192.168.1.1/repo.git") + + def test_hostname_allowed(self) -> None: + _validate_git_url("https://github.com/org/repo.git") class TestListReposOrgFilter: diff --git a/tests/unit/shared/test_http.py b/tests/unit/shared/test_http.py index 860ed6f3..b7e6774f 100644 --- a/tests/unit/shared/test_http.py +++ b/tests/unit/shared/test_http.py @@ -2,6 +2,7 @@ from __future__ import annotations +from contextlib import contextmanager from pathlib import Path from unittest.mock import MagicMock, patch @@ -18,6 +19,32 @@ ) +def _mock_stream_client(content: bytes, side_effect_error: Exception | None = None): + """Create a mock client whose .stream() returns a context manager yielding a response.""" + mock_resp = MagicMock(spec=httpx.Response) + if side_effect_error: + mock_resp.raise_for_status.side_effect = side_effect_error + else: + mock_resp.raise_for_status.return_value = None + + def iter_bytes(chunk_size: int = 65536): + for i in range(0, len(content), chunk_size): + yield content[i : i + chunk_size] + + mock_resp.iter_bytes = iter_bytes + + @contextmanager + def stream(method: str, url: str): + if side_effect_error is None: + yield mock_resp + else: + raise side_effect_error + + mock_client = MagicMock(spec=httpx.Client) + mock_client.stream.side_effect = stream + return mock_client + + class TestCreateClient: def test_returns_httpx_client(self) -> None: client = create_client() @@ -133,10 +160,7 @@ def test_private_url_not_skipped_when_check_ssrf_false(self) -> None: class TestDownloadFile: def test_successful_download(self, tmp_path: Path) -> None: output = tmp_path / "doc.md" - mock_resp = MagicMock(spec=httpx.Response) - mock_resp.content = b"hello world" - mock_client = MagicMock(spec=httpx.Client) - mock_client.get.return_value = mock_resp + mock_client = _mock_stream_client(b"hello world") with patch("src.shared.http.get_pooled_client", return_value=mock_client): ok, status = download_file("https://example.com/doc", output) @@ -156,18 +180,25 @@ def test_skips_private_url(self, tmp_path: Path) -> None: def test_retries_on_http_429(self, tmp_path: Path) -> None: output = tmp_path / "retry.md" + mock_resp_429 = MagicMock(spec=httpx.Response) mock_resp_429.status_code = 429 mock_resp_429.headers = {} - mock_resp_ok = MagicMock(spec=httpx.Response) - mock_resp_ok.content = b"content" + @contextmanager + def _cm_429(method: str, url: str): + yield mock_resp_429 + @contextmanager + def _cm_ok(method: str, url: str): + mock_resp_ok = MagicMock(spec=httpx.Response) + mock_resp_ok.raise_for_status.return_value = None + mock_resp_ok.iter_bytes = lambda chunk_size=65536: iter([b"content"]) + yield mock_resp_ok + + cm_sequence = iter([_cm_429, _cm_ok]) mock_client = MagicMock(spec=httpx.Client) - mock_client.get.side_effect = [ - httpx.HTTPStatusError("429", request=MagicMock(), response=mock_resp_429), - mock_resp_ok, - ] + mock_client.stream.side_effect = lambda m, u: next(cm_sequence)(m, u) with ( patch("src.shared.http.get_pooled_client", return_value=mock_client), @@ -179,11 +210,22 @@ def test_retries_on_http_429(self, tmp_path: Path) -> None: def test_retries_on_network_error(self, tmp_path: Path) -> None: output = tmp_path / "retry_net.md" + + @contextmanager + def _cm_fail(method: str, url: str): + raise httpx.ConnectError("timeout") + yield + + @contextmanager + def _cm_ok(method: str, url: str): + mock_resp = MagicMock(spec=httpx.Response) + mock_resp.raise_for_status.return_value = None + mock_resp.iter_bytes = lambda chunk_size=65536: iter([b"ok"]) + yield mock_resp + + cm_sequence = iter([_cm_fail, _cm_ok]) mock_client = MagicMock(spec=httpx.Client) - mock_client.get.side_effect = [ - httpx.ConnectError("timeout"), - MagicMock(content=b"ok"), - ] + mock_client.stream.side_effect = lambda m, u: next(cm_sequence)(m, u) with ( patch("src.shared.http.get_pooled_client", return_value=mock_client), @@ -195,8 +237,14 @@ def test_retries_on_network_error(self, tmp_path: Path) -> None: def test_gives_up_after_max_retries(self, tmp_path: Path) -> None: output = tmp_path / "fail.md" + + @contextmanager + def _cm_fail(method: str, url: str): + raise httpx.ConnectError("always fails") + yield + mock_client = MagicMock(spec=httpx.Client) - mock_client.get.side_effect = httpx.ConnectError("always fails") + mock_client.stream.side_effect = _cm_fail with ( patch("src.shared.http.get_pooled_client", return_value=mock_client), @@ -211,11 +259,17 @@ def test_non_retryable_http_status(self, tmp_path: Path) -> None: output = tmp_path / "forbidden.md" mock_resp = MagicMock(spec=httpx.Response) mock_resp.status_code = 403 - mock_client = MagicMock(spec=httpx.Client) - mock_client.get.side_effect = httpx.HTTPStatusError( + mock_resp.raise_for_status.side_effect = httpx.HTTPStatusError( "403", request=MagicMock(), response=mock_resp ) + @contextmanager + def _cm_forbidden(method: str, url: str): + yield mock_resp + + mock_client = MagicMock(spec=httpx.Client) + mock_client.stream.side_effect = _cm_forbidden + with patch("src.shared.http.get_pooled_client", return_value=mock_client): ok, status = download_file("https://example.com/forbidden", output) @@ -224,10 +278,7 @@ def test_non_retryable_http_status(self, tmp_path: Path) -> None: def test_filesystem_error(self, tmp_path: Path) -> None: output = Path("/nonexistent/path/doc.md") - mock_resp = MagicMock(spec=httpx.Response) - mock_resp.content = b"data" - mock_client = MagicMock(spec=httpx.Client) - mock_client.get.return_value = mock_resp + mock_client = _mock_stream_client(b"data") with patch("src.shared.http.get_pooled_client", return_value=mock_client): ok, status = download_file("https://example.com/doc", output) @@ -239,15 +290,24 @@ def test_respects_retry_after_header(self, tmp_path: Path) -> None: mock_resp_429 = MagicMock(spec=httpx.Response) mock_resp_429.status_code = 429 mock_resp_429.headers = {"Retry-After": "5"} + mock_resp_429.raise_for_status.side_effect = httpx.HTTPStatusError( + "429", request=MagicMock(), response=mock_resp_429 + ) + + @contextmanager + def _cm_429(method: str, url: str): + yield mock_resp_429 - mock_resp_ok = MagicMock(spec=httpx.Response) - mock_resp_ok.content = b"content" + @contextmanager + def _cm_ok(method: str, url: str): + mock_resp_ok = MagicMock(spec=httpx.Response) + mock_resp_ok.raise_for_status.return_value = None + mock_resp_ok.iter_bytes = lambda chunk_size=65536: iter([b"content"]) + yield mock_resp_ok + cm_sequence = iter([_cm_429, _cm_ok]) mock_client = MagicMock(spec=httpx.Client) - mock_client.get.side_effect = [ - httpx.HTTPStatusError("429", request=MagicMock(), response=mock_resp_429), - mock_resp_ok, - ] + mock_client.stream.side_effect = lambda m, u: next(cm_sequence)(m, u) with ( patch("src.shared.http.get_pooled_client", return_value=mock_client), diff --git a/tests/unit/workers/test_task_dispatcher.py b/tests/unit/workers/test_task_dispatcher.py index e9afbb89..68e29356 100644 --- a/tests/unit/workers/test_task_dispatcher.py +++ b/tests/unit/workers/test_task_dispatcher.py @@ -13,7 +13,7 @@ class TestTaskDispatcherInit: def test_default_poll_interval(self, mock_db: MagicMock) -> None: with patch("src.workers.processor.DatabaseService.get_instance", return_value=mock_db): dispatcher = TaskDispatcher() - assert dispatcher.poll_interval == 1.0 + assert dispatcher.poll_interval == 0.2 def test_custom_poll_interval(self, mock_db: MagicMock) -> None: with patch("src.workers.processor.DatabaseService.get_instance", return_value=mock_db): diff --git a/uv.lock b/uv.lock index bddb06f0..032c49bf 100644 --- a/uv.lock +++ b/uv.lock @@ -409,7 +409,7 @@ name = "cuda-bindings" version = "12.9.4" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "cuda-pathfinder", marker = "platform_machine != 's390x' and sys_platform != 'emscripten' and sys_platform != 'win32'" }, + { name = "cuda-pathfinder" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/a9/c1/dabe88f52c3e3760d861401bb994df08f672ec893b8f7592dc91626adcf3/cuda_bindings-12.9.4-cp312-cp312-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:fda147a344e8eaeca0c6ff113d2851ffca8f7dfc0a6c932374ee5c47caa649c8", size = 12151019, upload-time = "2025-10-21T14:51:43.167Z" }, @@ -1751,7 +1751,7 @@ name = "nvidia-cudnn-cu12" version = "9.10.2.21" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "nvidia-cublas-cu12", marker = "platform_machine != 's390x' and sys_platform != 'emscripten' and sys_platform != 'win32'" }, + { name = "nvidia-cublas-cu12" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/ba/51/e123d997aa098c61d029f76663dedbfb9bc8dcf8c60cbd6adbe42f76d049/nvidia_cudnn_cu12-9.10.2.21-py3-none-manylinux_2_27_x86_64.whl", hash = "sha256:949452be657fa16687d0930933f032835951ef0892b37d2d53824d1a84dc97a8", size = 706758467, upload-time = "2025-06-06T21:54:08.597Z" }, @@ -1762,7 +1762,7 @@ name = "nvidia-cufft-cu12" version = "11.3.3.83" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "nvidia-nvjitlink-cu12", marker = "platform_machine != 's390x' and sys_platform != 'emscripten' and sys_platform != 'win32'" }, + { name = "nvidia-nvjitlink-cu12" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/1f/13/ee4e00f30e676b66ae65b4f08cb5bcbb8392c03f54f2d5413ea99a5d1c80/nvidia_cufft_cu12-11.3.3.83-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:4d2dd21ec0b88cf61b62e6b43564355e5222e4a3fb394cac0db101f2dd0d4f74", size = 193118695, upload-time = "2025-03-07T01:45:27.821Z" }, @@ -1789,9 +1789,9 @@ name = "nvidia-cusolver-cu12" version = "11.7.3.90" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "nvidia-cublas-cu12", marker = "platform_machine != 's390x' and sys_platform != 'emscripten' and sys_platform != 'win32'" }, - { name = "nvidia-cusparse-cu12", marker = "platform_machine != 's390x' and sys_platform != 'emscripten' and sys_platform != 'win32'" }, - { name = "nvidia-nvjitlink-cu12", marker = "platform_machine != 's390x' and sys_platform != 'emscripten' and sys_platform != 'win32'" }, + { name = "nvidia-cublas-cu12" }, + { name = "nvidia-cusparse-cu12" }, + { name = "nvidia-nvjitlink-cu12" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/85/48/9a13d2975803e8cf2777d5ed57b87a0b6ca2cc795f9a4f59796a910bfb80/nvidia_cusolver_cu12-11.7.3.90-py3-none-manylinux_2_27_x86_64.whl", hash = "sha256:4376c11ad263152bd50ea295c05370360776f8c3427b30991df774f9fb26c450", size = 267506905, upload-time = "2025-03-07T01:47:16.273Z" }, @@ -1802,7 +1802,7 @@ name = "nvidia-cusparse-cu12" version = "12.5.8.93" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "nvidia-nvjitlink-cu12", marker = "platform_machine != 's390x' and sys_platform != 'emscripten' and sys_platform != 'win32'" }, + { name = "nvidia-nvjitlink-cu12" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/c2/f5/e1854cb2f2bcd4280c44736c93550cc300ff4b8c95ebe370d0aa7d2b473d/nvidia_cusparse_cu12-12.5.8.93-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:1ec05d76bbbd8b61b06a80e1eaf8cf4959c3d4ce8e711b65ebd0443bb0ebb13b", size = 288216466, upload-time = "2025-03-07T01:48:13.779Z" }, @@ -2768,6 +2768,7 @@ name = "sigmahqrag" version = "0.1.0" source = { virtual = "." } dependencies = [ + { name = "aiosqlite" }, { name = "docx2txt" }, { name = "duckdb" }, { name = "fastapi" }, @@ -2781,6 +2782,7 @@ dependencies = [ { name = "llama-index-llms-openai-like" }, { name = "llama-index-readers-file" }, { name = "llama-index-vector-stores-qdrant" }, + { name = "portalocker" }, { name = "puremagic" }, { name = "pymupdf" }, { name = "python-multipart" }, @@ -2805,6 +2807,7 @@ dev = [ [package.metadata] requires-dist = [ + { name = "aiosqlite", specifier = ">=0.22" }, { name = "docx2txt", specifier = ">=0.9" }, { name = "duckdb", specifier = ">=1.5.4" }, { name = "fastapi", specifier = ">=0.139.0" }, @@ -2818,6 +2821,7 @@ requires-dist = [ { name = "llama-index-llms-openai-like", specifier = ">=0.7.2" }, { name = "llama-index-readers-file", specifier = ">=0.6.0" }, { name = "llama-index-vector-stores-qdrant", specifier = ">=0.10.2" }, + { name = "portalocker", specifier = ">=3.2" }, { name = "puremagic", specifier = ">=2.2.0" }, { name = "pymupdf", specifier = ">=1.28.0" }, { name = "python-multipart", specifier = ">=0.0.32" },