diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index fcfd64a..00215eb 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -35,6 +35,8 @@ jobs: - name: Run pytest run: | uv run pytest tests/ -q \ + -m "not e2e" \ + --ignore tests/e2e \ --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 \ @@ -42,3 +44,13 @@ jobs: --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 + + e2e: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: astral-sh/setup-uv@v3 + - name: Install Playwright browser + run: uv run playwright install --with-deps chromium + - name: Run e2e tests + run: uv run pytest tests/e2e -q -m e2e diff --git a/pyproject.toml b/pyproject.toml index b2bd3fa..cdaa244 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -80,11 +80,16 @@ asyncio_mode = "auto" testpaths = ["tests"] addopts = "--import-mode=importlib" pythonpath = ["."] +markers = [ + "e2e: end-to-end UI tests that start a live server and browser", +] [dependency-groups] dev = [ "mypy", + "playwright>=1.48", "pre-commit>=4.6.0", + "pytest-playwright>=0.6", "pytest>=9.0.3", "pytest-asyncio>=1.4.0", "pytest-cov>=7.1.0", diff --git a/sigmaforge/__init__.py b/sigmaforge/__init__.py new file mode 100644 index 0000000..d3ec452 --- /dev/null +++ b/sigmaforge/__init__.py @@ -0,0 +1 @@ +__version__ = "0.2.0" diff --git a/sigmaforge/config.py b/sigmaforge/config.py new file mode 100644 index 0000000..bac95c9 --- /dev/null +++ b/sigmaforge/config.py @@ -0,0 +1,227 @@ +from __future__ import annotations + +import os +from dataclasses import asdict, dataclass, field +from pathlib import Path +from typing import TYPE_CHECKING, Any, Mapping + +if TYPE_CHECKING: + from sigmaforge.db.config import ConfigStore + +PROJECT_ROOT = Path(__file__).resolve().parent.parent +DATA_DIR = PROJECT_ROOT / "data" +DEFAULT_COLLECTIONS: tuple[str, ...] = ("sigma_rules", "sigma_docs", "sigma_spec") + + +def _default_data_dir() -> Path: + return DATA_DIR + + +def _default_duckdb_path() -> Path: + return DATA_DIR / "duckdb" / "sigmaforge.duckdb" + + +def _default_qdrant_storage_path() -> Path: + return DATA_DIR / "qdrant" + + +def _default_logs_dir() -> Path: + return DATA_DIR / "logs" + + +def _default_temp_dir() -> Path: + return DATA_DIR / "temp" + + +def _default_local_documents_dir() -> Path: + return DATA_DIR / "documents" / "local" + + +def _default_github_dir() -> Path: + return DATA_DIR / "github" + + +def _default_spec_dir() -> Path: + return DATA_DIR / "specification" + + +def _parse_bool(value: str) -> bool: + return value.strip().lower() in {"1", "true", "yes", "on"} + + +def _parse_int(value: str) -> int: + return int(value.strip()) + + +def _parse_float(value: str) -> float: + return float(value.strip()) + + +def _parse_collection(value: str) -> tuple[str, ...]: + parts = [part.strip() for part in value.split(",")] + return tuple(part for part in parts if part) + + +def _as_path(value: Any) -> Path: + return Path(str(value)) + + +def _as_str(value: Any) -> str: + return str(value) + + +def _as_int(value: Any) -> int: + return int(value) + + +def _as_float(value: Any) -> float: + return float(value) + + +def _as_bool(value: Any) -> bool: + if isinstance(value, bool): + return value + return _parse_bool(str(value)) + + +def _as_collections(value: Any) -> tuple[str, ...]: + if isinstance(value, str): + return _parse_collection(value) + return tuple(str(item) for item in value) + + +def _set_data_dir(config: Config, data_dir: Path) -> None: + config.data_dir = data_dir + config.duckdb_path = data_dir / "duckdb" / "sigmaforge.duckdb" + config.qdrant_storage_path = data_dir / "qdrant" + config.logs_dir = data_dir / "logs" + config.temp_dir = data_dir / "temp" + config.local_documents_dir = data_dir / "documents" / "local" + config.github_dir = data_dir / "github" + config.spec_dir = data_dir / "specification" + + +@dataclass(slots=True) +class Config: + data_dir: Path = field(default_factory=_default_data_dir) + duckdb_path: Path = field(default_factory=_default_duckdb_path) + qdrant_host: str = "127.0.0.1" + qdrant_port: int = 6333 + qdrant_storage_path: Path = field(default_factory=_default_qdrant_storage_path) + vector_size: int = 384 + collections: tuple[str, ...] = field(default_factory=lambda: DEFAULT_COLLECTIONS) + llama_base_url: str = "http://127.0.0.1:8080" + llama_model: str = "sigma" + llama_api_key: str = "sigma-key" + llama_timeout: float = 120.0 + embedding_model: str = "intfloat/multilingual-e5-small" + embedding_device: str = "cpu" + embedding_batch_size: int = 64 + logs_dir: Path = field(default_factory=_default_logs_dir) + temp_dir: Path = field(default_factory=_default_temp_dir) + local_documents_dir: Path = field(default_factory=_default_local_documents_dir) + github_dir: Path = field(default_factory=_default_github_dir) + spec_dir: Path = field(default_factory=_default_spec_dir) + chunk_size: int = 1024 + chunk_overlap: int = 100 + chunk_size_rules: int = 512 + chunk_overlap_rules: int = 50 + log_level: str = "INFO" + hf_offline: bool = True + + def apply_env(self, env: Mapping[str, str] | None = None) -> None: + source = os.environ if env is None else env + if value := source.get("SIGMAFORGE_DATA_DIR"): + _set_data_dir(self, Path(value)) + if value := source.get("SIGMAFORGE_DUCKDB_PATH"): + self.duckdb_path = Path(value) + if value := source.get("SIGMAFORGE_QDRANT_HOST"): + self.qdrant_host = value + if value := source.get("SIGMAFORGE_QDRANT_PORT"): + self.qdrant_port = _parse_int(value) + if value := source.get("SIGMAFORGE_QDRANT_STORAGE_PATH"): + self.qdrant_storage_path = Path(value) + if value := source.get("SIGMAFORGE_VECTOR_SIZE"): + self.vector_size = _parse_int(value) + if value := source.get("SIGMAFORGE_COLLECTIONS"): + self.collections = _parse_collection(value) + if value := source.get("SIGMAFORGE_LLAMA_BASE_URL"): + self.llama_base_url = value + if value := source.get("SIGMAFORGE_LLAMA_MODEL"): + self.llama_model = value + if value := source.get("SIGMAFORGE_LLAMA_API_KEY"): + self.llama_api_key = value + if value := source.get("SIGMAFORGE_LLAMA_TIMEOUT"): + self.llama_timeout = _parse_float(value) + if value := source.get("SIGMAFORGE_EMBEDDING_MODEL"): + self.embedding_model = value + if value := source.get("SIGMAFORGE_EMBEDDING_DEVICE"): + self.embedding_device = value + if value := source.get("SIGMAFORGE_EMBEDDING_BATCH_SIZE"): + self.embedding_batch_size = _parse_int(value) + if value := source.get("SIGMAFORGE_LOG_LEVEL"): + self.log_level = value + if value := source.get("HF_HUB_OFFLINE"): + self.hf_offline = _parse_bool(value) + + def apply_db(self, store: ConfigStore) -> None: + values = store.get_all() + if value := values.get("paths.data"): + _set_data_dir(self, _as_path(value)) + if value := values.get("paths.duckdb"): + self.duckdb_path = _as_path(value) + if value := values.get("services.qdrant.host"): + self.qdrant_host = _as_str(value) + if value := values.get("services.qdrant.port"): + self.qdrant_port = _as_int(value) + if value := values.get("services.qdrant.storage_path"): + self.qdrant_storage_path = _as_path(value) + if value := values.get("services.llama.base_url"): + self.llama_base_url = _as_str(value) + if value := values.get("services.llama.model"): + self.llama_model = _as_str(value) + if value := values.get("services.llama.api_key"): + self.llama_api_key = _as_str(value) + if value := values.get("models.embedding.active"): + self.embedding_model = _as_str(value) + if value := values.get("models.embedding.device"): + self.embedding_device = _as_str(value) + if value := values.get("models.embedding.batch_size"): + self.embedding_batch_size = _as_int(value) + if value := values.get("search.vector_size"): + self.vector_size = _as_int(value) + if value := values.get("search.collections"): + self.collections = _as_collections(value) + if value := values.get("logging.level"): + self.log_level = _as_str(value) + if value := values.get("ingestion.chunk_size"): + self.chunk_size = _as_int(value) + if value := values.get("ingestion.chunk_overlap"): + self.chunk_overlap = _as_int(value) + if value := values.get("ingestion.chunk_size_rules"): + self.chunk_size_rules = _as_int(value) + if value := values.get("ingestion.chunk_overlap_rules"): + self.chunk_overlap_rules = _as_int(value) + if value := values.get("ingestion.local_documents_dir"): + self.local_documents_dir = _as_path(value) + if value := values.get("ingestion.github_dir"): + self.github_dir = _as_path(value) + if value := values.get("ingestion.spec_dir"): + self.spec_dir = _as_path(value) + if value := values.get("hf_offline"): + self.hf_offline = _as_bool(value) + + def as_dict(self) -> dict[str, Any]: + return asdict(self) + + +def load_config( + *, + env: Mapping[str, str] | None = None, + config_store: ConfigStore | None = None, +) -> Config: + config = Config() + config.apply_env(env) + if config_store is not None: + config.apply_db(config_store) + return config diff --git a/sigmaforge/db/__init__.py b/sigmaforge/db/__init__.py new file mode 100644 index 0000000..26f9e7b --- /dev/null +++ b/sigmaforge/db/__init__.py @@ -0,0 +1,16 @@ +from sigmaforge.db.client import Database +from sigmaforge.db.config import ConfigStore +from sigmaforge.db.docs import DocRegistry, DocRecord +from sigmaforge.db.models import ModelRecord, ModelStore +from sigmaforge.db.tasks import Task, TaskStore + +__all__ = [ + "ConfigStore", + "Database", + "DocRecord", + "DocRegistry", + "ModelRecord", + "ModelStore", + "Task", + "TaskStore", +] diff --git a/sigmaforge/db/client.py b/sigmaforge/db/client.py new file mode 100644 index 0000000..1b56e7b --- /dev/null +++ b/sigmaforge/db/client.py @@ -0,0 +1,101 @@ +from __future__ import annotations + +import threading +from pathlib import Path +from typing import Any, Sequence + +import duckdb + +from sigmaforge.errors import DatabaseError + +SCHEMA_STATEMENTS: tuple[str, ...] = ( + """ + CREATE TABLE IF NOT EXISTS config ( + key TEXT PRIMARY KEY, + value TEXT NOT NULL, + updated_at TIMESTAMP DEFAULT now() + ) + """, + """ + CREATE TABLE IF NOT EXISTS models ( + kind TEXT NOT NULL, + name TEXT NOT NULL, + path TEXT NOT NULL, + active BOOLEAN DEFAULT FALSE, + created_at TIMESTAMP DEFAULT now(), + PRIMARY KEY (kind, name) + ) + """, + """ + CREATE TABLE IF NOT EXISTS tasks ( + id TEXT PRIMARY KEY, + type TEXT NOT NULL, + status TEXT NOT NULL, + progress DOUBLE NOT NULL DEFAULT 0.0, + message TEXT NOT NULL DEFAULT '', + payload TEXT NOT NULL DEFAULT '{}', + created_at TIMESTAMP DEFAULT now(), + updated_at TIMESTAMP DEFAULT now() + ) + """, + """ + CREATE TABLE IF NOT EXISTS docs_registry ( + source TEXT NOT NULL, + path TEXT NOT NULL, + status TEXT NOT NULL, + indexed_at TIMESTAMP, + PRIMARY KEY (source, path) + ) + """, +) + + +class Database: + def __init__(self, path: str | Path = ":memory:") -> None: + self.path = Path(path) if path != ":memory:" else Path(":memory:") + self._lock = threading.RLock() + try: + self._conn = duckdb.connect(str(path)) + except duckdb.Error as exc: + raise DatabaseError(str(exc)) from exc + + @property + def connection(self) -> duckdb.DuckDBPyConnection: + return self._conn + + def execute(self, sql: str, params: Sequence[Any] = ()) -> None: + with self._lock: + try: + self._conn.execute(sql, list(params)) + self._conn.commit() + except duckdb.Error as exc: + raise DatabaseError(str(exc)) from exc + + def fetch_all(self, sql: str, params: Sequence[Any] = ()) -> list[tuple[Any, ...]]: + with self._lock: + try: + cursor = self._conn.execute(sql, list(params)) + return cursor.fetchall() + except duckdb.Error as exc: + raise DatabaseError(str(exc)) from exc + + def fetch_one(self, sql: str, params: Sequence[Any] = ()) -> tuple[Any, ...] | None: + rows = self.fetch_all(sql, params) + return rows[0] if rows else None + + def fetch_dict(self, sql: str, params: Sequence[Any] = ()) -> dict[str, Any] | None: + row = self.fetch_one(sql, params) + if row is None: + return None + cursor = self._conn.description + if cursor is None: + return dict(zip(("value",), row)) + return {column[0]: value for column, value in zip(cursor, row)} + + def init_schema(self) -> None: + for statement in SCHEMA_STATEMENTS: + self.execute(statement) + + def close(self) -> None: + with self._lock: + self._conn.close() diff --git a/sigmaforge/db/config.py b/sigmaforge/db/config.py new file mode 100644 index 0000000..d64f2a9 --- /dev/null +++ b/sigmaforge/db/config.py @@ -0,0 +1,51 @@ +from __future__ import annotations + +import json +from typing import Any + +from sigmaforge.db.client import Database + + +def _encode(value: Any) -> str: + return json.dumps(value) + + +def _decode(value: str) -> Any: + try: + return json.loads(value) + except json.JSONDecodeError: + return value + + +class ConfigStore: + def __init__(self, db: Database) -> None: + self.db = db + + def get(self, key: str, default: Any = None) -> Any: + row = self.db.fetch_one("SELECT value FROM config WHERE key = ?", [key]) + if row is None: + return default + return _decode(row[0]) + + def set(self, key: str, value: Any) -> None: + self.db.execute( + """ + INSERT INTO config (key, value, updated_at) + VALUES (?, ?, now()) + ON CONFLICT (key) DO UPDATE SET + value = excluded.value, + updated_at = now() + """, + [key, _encode(value)], + ) + + def delete(self, key: str) -> bool: + row = self.db.fetch_one("SELECT 1 FROM config WHERE key = ?", [key]) + if row is None: + return False + self.db.execute("DELETE FROM config WHERE key = ?", [key]) + return True + + def get_all(self) -> dict[str, Any]: + rows = self.db.fetch_all("SELECT key, value FROM config") + return {key: _decode(value) for key, value in rows} diff --git a/sigmaforge/db/docs.py b/sigmaforge/db/docs.py new file mode 100644 index 0000000..5214b36 --- /dev/null +++ b/sigmaforge/db/docs.py @@ -0,0 +1,100 @@ +from __future__ import annotations + +from dataclasses import dataclass +from datetime import datetime +from typing import Any + +from sigmaforge.db.client import Database + + +@dataclass(frozen=True, slots=True) +class DocRecord: + source: str + path: str + status: str + indexed_at: str | None = None + + +def _format_timestamp(value: Any) -> str | None: + if value is None: + return None + if isinstance(value, datetime): + return value.isoformat() + return str(value) + + +class DocRegistry: + def __init__(self, db: Database) -> None: + self.db = db + + def mark( + self, + source: str, + path: str, + status: str, + indexed_at: str | None = None, + ) -> None: + self.db.execute( + """ + INSERT INTO docs_registry (source, path, status, indexed_at) + VALUES (?, ?, ?, ?) + ON CONFLICT (source, path) DO UPDATE SET + status = excluded.status, + indexed_at = excluded.indexed_at + """, + [source, path, status, indexed_at], + ) + + def get(self, source: str, path: str) -> DocRecord | None: + row = self.db.fetch_one( + "SELECT source, path, status, indexed_at FROM docs_registry WHERE source = ? AND path = ?", + [source, path], + ) + if row is None: + return None + return DocRecord( + source=row[0], + path=row[1], + status=row[2], + indexed_at=_format_timestamp(row[3]), + ) + + def list_records(self, status: str | None = None) -> list[DocRecord]: + if status is None: + rows = self.db.fetch_all( + "SELECT source, path, status, indexed_at FROM docs_registry ORDER BY source, path" + ) + else: + rows = self.db.fetch_all( + """ + SELECT source, path, status, indexed_at + FROM docs_registry + WHERE status = ? + ORDER BY source, path + """, + [status], + ) + return [ + DocRecord( + source=row[0], + path=row[1], + status=row[2], + indexed_at=_format_timestamp(row[3]), + ) + for row in rows + ] + + def delete(self, source: str, path: str) -> bool: + if self.get(source, path) is None: + return False + self.db.execute( + "DELETE FROM docs_registry WHERE source = ? AND path = ?", + [source, path], + ) + return True + + def to_dicts(self) -> list[dict[str, Any]]: + return [ + {"source": item.source, "path": item.path, "status": item.status, "indexed_at": item.indexed_at} + for item in self.list_records() + ] diff --git a/sigmaforge/db/models.py b/sigmaforge/db/models.py new file mode 100644 index 0000000..2f5fd6a --- /dev/null +++ b/sigmaforge/db/models.py @@ -0,0 +1,77 @@ +from __future__ import annotations + +from dataclasses import dataclass + +from sigmaforge.db.client import Database + + +@dataclass(frozen=True, slots=True) +class ModelRecord: + kind: str + name: str + path: str + active: bool + + +class ModelStore: + def __init__(self, db: Database) -> None: + self.db = db + + def upsert(self, kind: str, name: str, path: str, active: bool = False) -> None: + self.db.execute( + """ + INSERT INTO models (kind, name, path, active) + VALUES (?, ?, ?, ?) + ON CONFLICT (kind, name) DO UPDATE SET + path = excluded.path, + active = excluded.active + """, + [kind, name, path, active], + ) + + def get(self, kind: str, name: str) -> ModelRecord | None: + row = self.db.fetch_one( + "SELECT kind, name, path, active FROM models WHERE kind = ? AND name = ?", + [kind, name], + ) + if row is None: + return None + return ModelRecord(kind=row[0], name=row[1], path=row[2], active=bool(row[3])) + + def get_active(self, kind: str) -> str | None: + row = self.db.fetch_one( + "SELECT name FROM models WHERE kind = ? AND active = TRUE ORDER BY name LIMIT 1", + [kind], + ) + return row[0] if row else None + + def set_active(self, kind: str, name: str) -> bool: + if self.get(kind, name) is None: + return False + self.db.execute( + "UPDATE models SET active = FALSE WHERE kind = ? AND name != ?", + [kind, name], + ) + self.db.execute( + "UPDATE models SET active = TRUE WHERE kind = ? AND name = ?", + [kind, name], + ) + return True + + def list(self, kind: str | None = None) -> list[ModelRecord]: + if kind is None: + rows = self.db.fetch_all( + "SELECT kind, name, path, active FROM models ORDER BY kind, name" + ) + else: + rows = self.db.fetch_all( + "SELECT kind, name, path, active FROM models WHERE kind = ? ORDER BY name", + [kind], + ) + return [ModelRecord(row[0], row[1], row[2], bool(row[3])) for row in rows] + + def delete(self, kind: str, name: str) -> bool: + if self.get(kind, name) is None: + return False + self.db.execute("DELETE FROM models WHERE kind = ? AND name = ?", [kind, name]) + return True diff --git a/sigmaforge/db/tasks.py b/sigmaforge/db/tasks.py new file mode 100644 index 0000000..d82f085 --- /dev/null +++ b/sigmaforge/db/tasks.py @@ -0,0 +1,124 @@ +from __future__ import annotations + +import json +from dataclasses import dataclass, field +from datetime import datetime +from typing import Any + +from sigmaforge.db.client import Database + + +def _timestamp(value: Any) -> str | None: + if isinstance(value, datetime): + return value.isoformat() + return None if value is None else str(value) + + +@dataclass(slots=True) +class Task: + id: str + type: str + status: str + progress: float = 0.0 + message: str = "" + payload: dict[str, Any] = field(default_factory=dict) + created_at: str | None = None + updated_at: str | None = None + + +class TaskStore: + def __init__(self, db: Database) -> None: + self.db = db + + def create(self, task: Task) -> None: + self.db.execute( + """ + INSERT INTO tasks (id, type, status, progress, message, payload) + VALUES (?, ?, ?, ?, ?, ?) + """, + [ + task.id, + task.type, + task.status, + task.progress, + task.message, + json.dumps(task.payload), + ], + ) + + def get(self, task_id: str) -> Task | None: + row = self.db.fetch_one( + """ + SELECT id, type, status, progress, message, payload, created_at, updated_at + FROM tasks + WHERE id = ? + """, + [task_id], + ) + if row is None: + return None + return Task( + id=row[0], + type=row[1], + status=row[2], + progress=float(row[3]), + message=row[4], + payload=json.loads(row[5]), + created_at=_timestamp(row[6]), + updated_at=_timestamp(row[7]), + ) + + def update_status( + self, + task_id: str, + status: str, + progress: float | None = None, + message: str | None = None, + ) -> bool: + current = self.get(task_id) + if current is None: + return False + next_progress = current.progress if progress is None else progress + next_message = current.message if message is None else message + self.db.execute( + """ + UPDATE tasks + SET status = ?, progress = ?, message = ?, updated_at = now() + WHERE id = ? + """, + [status, next_progress, next_message, task_id], + ) + return True + + def update_payload(self, task_id: str, payload: dict[str, Any]) -> bool: + if self.get(task_id) is None: + return False + self.db.execute( + "UPDATE tasks SET payload = ?, updated_at = current_timestamp WHERE id = ?", + [json.dumps(payload), task_id], + ) + return True + + def list(self, limit: int = 100) -> list[Task]: + rows = self.db.fetch_all( + """ + SELECT id, type, status, progress, message, payload, created_at, updated_at + FROM tasks + ORDER BY created_at DESC + LIMIT ? + """, + [limit], + ) + return [ + Task( + id=row[0], + type=row[1], + status=row[2], + progress=float(row[3]), + message=row[4], + payload=json.loads(row[5]), + created_at=_timestamp(row[6]), + updated_at=_timestamp(row[7]), + ) + for row in rows + ] diff --git a/sigmaforge/embed/__init__.py b/sigmaforge/embed/__init__.py new file mode 100644 index 0000000..ad70196 --- /dev/null +++ b/sigmaforge/embed/__init__.py @@ -0,0 +1,9 @@ +from sigmaforge.embed.dense import DenseEncoder +from sigmaforge.embed.sparse import Bm25SparseEncoder, SparseVector, create_sparse_encoder + +__all__ = [ + "Bm25SparseEncoder", + "DenseEncoder", + "SparseVector", + "create_sparse_encoder", +] diff --git a/sigmaforge/embed/dense.py b/sigmaforge/embed/dense.py new file mode 100644 index 0000000..0a090ba --- /dev/null +++ b/sigmaforge/embed/dense.py @@ -0,0 +1,60 @@ +from __future__ import annotations + +from pathlib import Path +from typing import Any, Sequence + + +class DenseEncoder: + def __init__( + self, + model_name: str, + model_dir: str | Path | None = None, + device: str = "cpu", + batch_size: int = 64, + query_prefix: str = "", + passage_prefix: str = "", + model: Any | None = None, + ) -> None: + self.model_name = model_name + self.model_dir = Path(model_dir) if model_dir is not None else None + self.device = device + self.batch_size = batch_size + self.query_prefix = query_prefix + self.passage_prefix = passage_prefix + self._model = model + + def _resolve_model_name(self) -> str: + if self.model_dir is not None: + candidate = self.model_dir / self.model_name + if candidate.exists(): + return str(candidate) + return self.model_name + + def _get_model(self) -> Any: + if self._model is None: + from sentence_transformers import SentenceTransformer + + self._model = SentenceTransformer( + self._resolve_model_name(), + device=self.device, + ) + return self._model + + def encode(self, texts: Sequence[str], *, is_query: bool = False) -> list[list[float]]: + if not texts: + return [] + model = self._get_model() + prefix = self.query_prefix if is_query else self.passage_prefix + prefixed = [f"{prefix}{text}" for text in texts] + vectors = model.encode( + prefixed, + batch_size=self.batch_size, + normalize_embeddings=True, + show_progress_bar=False, + ) + return [list(map(float, vector)) for vector in vectors] + + @property + def dim(self) -> int: + model = self._get_model() + return int(model.get_sentence_embedding_dimension()) diff --git a/sigmaforge/embed/sparse.py b/sigmaforge/embed/sparse.py new file mode 100644 index 0000000..f326691 --- /dev/null +++ b/sigmaforge/embed/sparse.py @@ -0,0 +1,133 @@ +from __future__ import annotations + +import hashlib +import math +import re +from dataclasses import dataclass +from typing import Sequence + +STOP_WORDS: frozenset[str] = frozenset( + { + "what", + "are", + "the", + "in", + "for", + "is", + "of", + "and", + "to", + "how", + "does", + "a", + "an", + "at", + "on", + "with", + "as", + "not", + "be", + "or", + "from", + "by", + "it", + "its", + "that", + "this", + "which", + "can", + "when", + "if", + "where", + "will", + "use", + "used", + "vs", + "between", + "than", + "but", + "must", + "should", + "would", + "shall", + "may", + "might", + "need", + "do", + "did", + "has", + "have", + "had", + "being", + "been", + "were", + "was", + "into", + "about", + "up", + "out", + "all", + "any", + "each", + "every", + "both", + "few", + "more", + "most", + "other", + "some", + "such", + "only", + "own", + "same", + "so", + "too", + "very", + "just", + "also", + "then", + "now", + "no", + "yes", + "q", + } +) + +_TOKEN_PATTERN = re.compile(r"\b[a-zA-Z][a-zA-Z0-9_]{2,}\b") +_MAX_HASH_ID = 2**24 + + +def _token_id(token: str) -> int: + return int(hashlib.md5(token.encode()).hexdigest()[:8], 16) % _MAX_HASH_ID + + +@dataclass(frozen=True, slots=True) +class SparseVector: + indices: tuple[int, ...] + values: tuple[float, ...] + + +class Bm25SparseEncoder: + def __call__(self, texts: Sequence[str]) -> list[SparseVector]: + return self.encode(texts) + + def encode(self, texts: Sequence[str]) -> list[SparseVector]: + return [self.encode_text(text) for text in texts] + + def encode_text(self, text: str) -> SparseVector: + tokens = [token for token in _TOKEN_PATTERN.findall(text.lower()) if token not in STOP_WORDS] + if not tokens: + return SparseVector((), ()) + frequencies: dict[str, int] = {} + for token in tokens: + frequencies[token] = frequencies.get(token, 0) + 1 + indices: list[int] = [] + values: list[float] = [] + for token, frequency in frequencies.items(): + values.append(1.0 + math.log(frequency)) + indices.append(_token_id(token)) + return SparseVector(indices=tuple(indices), values=tuple(values)) + + +def create_sparse_encoder() -> Bm25SparseEncoder: + return Bm25SparseEncoder() diff --git a/sigmaforge/errors.py b/sigmaforge/errors.py new file mode 100644 index 0000000..004983f --- /dev/null +++ b/sigmaforge/errors.py @@ -0,0 +1,29 @@ +from __future__ import annotations + + +class SigmaForgeError(Exception): + pass + + +class ConfigError(SigmaForgeError): + pass + + +class DatabaseError(SigmaForgeError): + pass + + +class QdrantError(SigmaForgeError): + pass + + +class EmbeddingError(SigmaForgeError): + pass + + +class LlamaError(SigmaForgeError): + pass + + +class SearchError(SigmaForgeError): + pass diff --git a/sigmaforge/ingestion/__init__.py b/sigmaforge/ingestion/__init__.py new file mode 100644 index 0000000..a09f354 --- /dev/null +++ b/sigmaforge/ingestion/__init__.py @@ -0,0 +1,28 @@ +from sigmaforge.ingestion.models import Chunk, IngestRequest, IngestResult, IngestSummary +from sigmaforge.ingestion.pipeline import IngestionPipeline +from sigmaforge.ingestion.sigma import ( + SigmaRule, + chunk_rule, + flat_rule_text, + iter_sigma_rule_files, + parse_sigma_rule, + rich_rule_chunks, + rule_metadata, +) +from sigmaforge.ingestion.text import split_text + +__all__ = [ + "Chunk", + "IngestRequest", + "IngestResult", + "IngestSummary", + "IngestionPipeline", + "SigmaRule", + "chunk_rule", + "flat_rule_text", + "iter_sigma_rule_files", + "parse_sigma_rule", + "rich_rule_chunks", + "rule_metadata", + "split_text", +] diff --git a/sigmaforge/ingestion/files.py b/sigmaforge/ingestion/files.py new file mode 100644 index 0000000..1eb521f --- /dev/null +++ b/sigmaforge/ingestion/files.py @@ -0,0 +1,47 @@ +from __future__ import annotations + +from collections.abc import Iterable, Sequence +from pathlib import Path + + +def iter_files( + directory: str | Path, + extensions: Sequence[str] = (), + recursive: bool = True, + selected_dirs: Sequence[str] = (), +) -> list[Path]: + root = Path(directory) + if not root.exists() or not root.is_dir(): + return [] + normalized = {str(extension).lower() for extension in extensions} + pattern = "**/*" if recursive else "*" + files: list[Path] = [] + for path in root.glob(pattern): + if not path.is_file(): + continue + if normalized and path.suffix.lower() not in normalized: + continue + if selected_dirs and not _matches_selected_dirs(path, root, selected_dirs): + continue + files.append(path) + return sorted(files) + + +def _matches_selected_dirs(path: Path, root: Path, selected_dirs: Iterable[str]) -> bool: + try: + rel = path.relative_to(root) + except ValueError: + return False + rel_parts = rel.parts + parent_parts = rel.parent.parts + for raw in selected_dirs: + parts = Path(str(raw)).parts + if not parts: + continue + if parts[-1] in parent_parts: + return True + if parent_parts[: len(parts)] == parts: + return True + if rel_parts[: len(parts)] == parts: + return True + return False diff --git a/sigmaforge/ingestion/models.py b/sigmaforge/ingestion/models.py new file mode 100644 index 0000000..a32a021 --- /dev/null +++ b/sigmaforge/ingestion/models.py @@ -0,0 +1,88 @@ +from __future__ import annotations + +import uuid +from dataclasses import dataclass, field +from typing import Any + + +@dataclass(slots=True) +class IngestRequest: + directory: str | None = None + recursive: bool = True + mode: str = "flat" + selected_dirs: list[str] = field(default_factory=list) + + +@dataclass(slots=True) +class IngestResult: + file: str + success: bool + rule_id: str | None = None + error: str | None = None + chunks: int = 0 + + def to_dict(self) -> dict[str, Any]: + return { + "file": self.file, + "success": self.success, + "rule_id": self.rule_id, + "error": self.error, + "chunks": self.chunks, + } + + +@dataclass(slots=True) +class IngestSummary: + collection: str + total: int = 0 + indexed: int = 0 + failed: int = 0 + chunks: int = 0 + results: list[IngestResult] = field(default_factory=list) + + @classmethod + def from_results( + cls, + collection: str, + results: list[IngestResult], + chunks: int = 0, + ) -> IngestSummary: + success = sum(1 for result in results if result.success) + return cls( + collection=collection, + total=len(results), + indexed=success, + failed=len(results) - success, + chunks=chunks, + results=list(results), + ) + + def to_dict(self) -> dict[str, Any]: + return { + "collection": self.collection, + "total": self.total, + "indexed": self.indexed, + "failed": self.failed, + "chunks": self.chunks, + "results": [result.to_dict() for result in self.results], + } + + +@dataclass(frozen=True, slots=True) +class Chunk: + text: str + source_file: str + chunk_type: str + chunk_index: int = 0 + metadata: dict[str, Any] = field(default_factory=dict) + + def payload(self) -> dict[str, Any]: + payload = dict(self.metadata) + payload["source_file"] = self.source_file + payload["chunk_type"] = self.chunk_type + payload["chunk_index"] = self.chunk_index + return payload + + def point_id(self, collection: str) -> str: + identifier = f"{collection}:{self.source_file}:{self.chunk_type}:{self.chunk_index}" + return str(uuid.uuid5(uuid.NAMESPACE_URL, identifier)) diff --git a/sigmaforge/ingestion/pipeline.py b/sigmaforge/ingestion/pipeline.py new file mode 100644 index 0000000..8d35d5d --- /dev/null +++ b/sigmaforge/ingestion/pipeline.py @@ -0,0 +1,207 @@ +from __future__ import annotations + +import logging +from collections.abc import Sequence +from pathlib import Path + +from sigmaforge.db.docs import DocRegistry +from sigmaforge.embed.dense import DenseEncoder +from sigmaforge.embed.sparse import Bm25SparseEncoder, SparseVector +from sigmaforge.ingestion.files import iter_files +from sigmaforge.ingestion.models import Chunk, IngestRequest, IngestResult, IngestSummary +from sigmaforge.ingestion.sigma import ( + SigmaRule, + chunk_rule, + iter_sigma_rule_files, + parse_sigma_rule, +) +from sigmaforge.ingestion.text import split_text +from sigmaforge.qdrant.store import QdrantStore + +logger = logging.getLogger(__name__) + +DEFAULT_TEXT_EXTENSIONS: tuple[str, ...] = (".md", ".markdown", ".rst", ".txt", ".adoc") + + +class IngestionPipeline: + def __init__( + self, + store: QdrantStore, + dense_encoder: DenseEncoder, + sparse_encoder: Bm25SparseEncoder | None = None, + registry: DocRegistry | None = None, + *, + default_directory: str | Path | None = None, + chunk_size: int = 1024, + chunk_overlap: int = 100, + enable_hybrid: bool | None = None, + ) -> None: + self.store = store + self.dense_encoder = dense_encoder + self.sparse_encoder = sparse_encoder + self.registry = registry + self.default_directory = default_directory + self.chunk_size = chunk_size + self.chunk_overlap = chunk_overlap + self.enable_hybrid = ( + store.manager.enable_hybrid if enable_hybrid is None else bool(enable_hybrid) + ) + + def index_chunks(self, collection: str, chunks: Sequence[Chunk]) -> int: + prepared = [chunk for chunk in chunks if chunk.text.strip()] + if not prepared: + return 0 + self.store.manager.ensure_collection(collection) + grouped: dict[str, list[Chunk]] = {} + for chunk in prepared: + grouped.setdefault(chunk.source_file, []).append(chunk) + total = 0 + for source_file, source_chunks in grouped.items(): + texts = [chunk.text for chunk in source_chunks] + dense_vectors = self.dense_encoder.encode(texts) + sparse_vectors = self._encode_sparse(texts) + payloads = [chunk.payload() for chunk in source_chunks] + ids = [chunk.point_id(collection) for chunk in source_chunks] + self.store.replace_source_points( + collection=collection, + source_file=source_file, + texts=texts, + dense_vectors=dense_vectors, + payloads=payloads, + sparse_vectors=sparse_vectors, + ids=ids, + ) + total += len(texts) + return total + + def ingest_sigma_directory( + self, + request: IngestRequest | None = None, + collection: str = "sigma_rules", + ) -> IngestSummary: + request = request or IngestRequest() + if request.mode not in {"flat", "rich"}: + raise ValueError("mode must be 'flat' or 'rich'") + directory = self._resolve_directory(request.directory) + files = iter_sigma_rule_files( + directory, + recursive=request.recursive, + selected_dirs=request.selected_dirs, + ) + results: list[IngestResult] = [] + chunks: list[Chunk] = [] + for path in files: + result = IngestResult(file=str(path), success=False) + try: + rule = parse_sigma_rule(path) + if rule is None: + raise ValueError("not a valid sigma rule") + created = chunk_rule(rule, request.mode) + result.success = True + result.rule_id = rule.id + result.chunks = len(created) + chunks.extend(created) + except Exception as exc: + result.error = str(exc) + logger.warning("Failed to ingest sigma rule %s: %s", path, exc) + results.append(result) + self._mark(str(path), collection, result.success) + indexed_chunks = self.index_chunks(collection, chunks) + return IngestSummary.from_results(collection, results, chunks=indexed_chunks) + + def ingest_sigma_rules( + self, + rules: Sequence[SigmaRule], + collection: str = "sigma_rules", + mode: str = "flat", + ) -> IngestSummary: + if mode not in {"flat", "rich"}: + raise ValueError("mode must be 'flat' or 'rich'") + results: list[IngestResult] = [] + chunks: list[Chunk] = [] + for rule in rules: + source = rule.file_path or rule.id + result = IngestResult(file=source, success=False, rule_id=rule.id) + try: + created = chunk_rule(rule, mode) + result.success = True + result.chunks = len(created) + chunks.extend(created) + except Exception as exc: + result.error = str(exc) + logger.warning("Failed to ingest sigma rule %s: %s", source, exc) + results.append(result) + self._mark(source, collection, result.success) + indexed_chunks = self.index_chunks(collection, chunks) + return IngestSummary.from_results(collection, results, chunks=indexed_chunks) + + def ingest_text_directory( + self, + request: IngestRequest | None = None, + collection: str = "sigma_docs", + extensions: Sequence[str] = DEFAULT_TEXT_EXTENSIONS, + source: str = "local", + ) -> IngestSummary: + request = request or IngestRequest() + directory = self._resolve_directory(request.directory) + files = iter_files( + directory, + extensions=extensions, + recursive=request.recursive, + selected_dirs=request.selected_dirs, + ) + results: list[IngestResult] = [] + chunks: list[Chunk] = [] + for path in files: + result = IngestResult(file=str(path), success=False) + try: + text = path.read_text(encoding="utf-8", errors="replace") + pieces = split_text(text, self.chunk_size, self.chunk_overlap) + result.success = True + result.chunks = len(pieces) + for index, piece in enumerate(pieces): + chunks.append( + Chunk( + text=piece, + source_file=str(path), + chunk_type="doc", + chunk_index=index, + metadata={"source": source, "title": path.stem}, + ) + ) + except Exception as exc: + result.error = str(exc) + logger.warning("Failed to ingest text file %s: %s", path, exc) + results.append(result) + self._mark(str(path), source, result.success) + indexed_chunks = self.index_chunks(collection, chunks) + return IngestSummary.from_results(collection, results, chunks=indexed_chunks) + + def _encode_sparse( + self, texts: Sequence[str] + ) -> Sequence[tuple[Sequence[int], Sequence[float]]] | None: + if self.sparse_encoder is None or not self.enable_hybrid: + return None + vectors = self.sparse_encoder(texts) + return [self._normalize_sparse(vector) for vector in vectors] + + def _mark(self, path: str, source: str, success: bool) -> None: + if self.registry is None: + return + self.registry.mark(source, path, "indexed" if success else "failed") + + def _resolve_directory(self, directory: str | None) -> Path: + candidate = Path(directory or self.default_directory or "") + if not candidate: + raise FileNotFoundError("No ingestion directory provided") + if not candidate.exists() or not candidate.is_dir(): + raise FileNotFoundError(f"Ingestion directory not found: {candidate}") + return candidate + + @staticmethod + def _normalize_sparse( + vector: SparseVector | tuple[Sequence[int], Sequence[float]], + ) -> tuple[Sequence[int], Sequence[float]]: + if isinstance(vector, tuple): + return vector[0], vector[1] + return vector.indices, vector.values diff --git a/sigmaforge/ingestion/sigma.py b/sigmaforge/ingestion/sigma.py new file mode 100644 index 0000000..3c1884e --- /dev/null +++ b/sigmaforge/ingestion/sigma.py @@ -0,0 +1,313 @@ +from __future__ import annotations + +from collections.abc import Mapping, Sequence +from dataclasses import asdict, dataclass, field +from datetime import date, datetime +from pathlib import Path +from typing import Any + +import yaml + +from sigmaforge.ingestion.files import iter_files +from sigmaforge.ingestion.models import Chunk + +_YAML_EXTENSIONS: tuple[str, ...] = (".yaml", ".yml") +_MAX_RULE_BYTES: int = 1024 * 1024 + + +@dataclass(slots=True) +class SigmaRule: + id: str + title: str + detection: dict[str, Any] = field(default_factory=dict) + condition: str = "" + status: str | None = None + level: str | None = None + tags: list[str] = field(default_factory=list) + falsepositives: list[str] = field(default_factory=list) + description: str | None = None + author: str | None = None + date: str | None = None + modified: str | None = None + references: list[str] = field(default_factory=list) + logsource: dict[str, Any] = field(default_factory=dict) + license: str | None = None + file_path: str | None = None + line_number: int | None = None + fields: list[str] = field(default_factory=list) + + def to_dict(self) -> dict[str, Any]: + return asdict(self) + + +def is_sigma_rule_dict(data: object) -> bool: + return isinstance(data, dict) and isinstance(data.get("detection"), dict) + + +def iter_sigma_rule_files( + directory: str | Path, + recursive: bool = True, + selected_dirs: Sequence[str] = (), +) -> list[Path]: + return iter_files( + directory, + extensions=_YAML_EXTENSIONS, + recursive=recursive, + selected_dirs=selected_dirs, + ) + + +def parse_sigma_rule(path: str | Path) -> SigmaRule | None: + path = Path(path) + if path.suffix.lower() not in _YAML_EXTENSIONS or not path.is_file(): + return None + try: + if path.stat().st_size > _MAX_RULE_BYTES: + return None + data = yaml.safe_load(path.read_text(encoding="utf-8")) + except (OSError, UnicodeDecodeError, yaml.YAMLError): + return None + if not is_sigma_rule_dict(data): + return None + detection = _as_mapping(data.get("detection")) + return SigmaRule( + id=str(data.get("id") or path.stem), + title=str(data.get("title") or path.stem), + detection=detection, + condition=str(data.get("condition") or detection.get("condition") or ""), + status=_as_optional_str(data.get("status")), + level=_as_optional_str(data.get("level")), + tags=_as_str_list(data.get("tags")), + falsepositives=_as_str_list(data.get("falsepositives")), + description=_as_optional_str(data.get("description")), + author=_as_optional_str(data.get("author")), + date=_normalize_date(data.get("date")), + modified=_normalize_date(data.get("modified")), + references=_as_str_list(data.get("references")), + logsource=_as_mapping(data.get("logsource")), + license=_as_optional_str(data.get("license")), + file_path=str(path), + line_number=_as_optional_int(data.get("line_number")), + fields=_as_str_list(data.get("fields")), + ) + + +def flat_rule_text(rule: SigmaRule) -> str: + lines = [f"Title: {rule.title}", f"Rule ID: {rule.id}"] + if rule.status: + lines.append(f"Status: {rule.status}") + if rule.level: + lines.append(f"Level: {rule.level}") + if rule.description: + lines.append(f"Description: {rule.description}") + if rule.author: + lines.append(f"Author: {rule.author}") + if rule.date: + lines.append(f"Date: {rule.date}") + if rule.modified: + lines.append(f"Modified: {rule.modified}") + if rule.tags: + lines.append(f"Tags: {', '.join(rule.tags)}") + if rule.references: + lines.append(f"References: {', '.join(rule.references)}") + if rule.logsource: + lines.append(f"Logsource: {format_logsource(rule.logsource)}") + if rule.condition: + lines.append(f"Condition: {rule.condition}") + detection_lines = [ + f" {key}: {_format_value(value)}" + for key, value in rule.detection.items() + if key != "condition" and _format_value(value) + ] + if detection_lines: + lines.append("Detection:") + lines.extend(detection_lines) + if rule.falsepositives: + lines.append(f"False positives: {'; '.join(rule.falsepositives)}") + return "\n".join(lines) + + +def rule_metadata(rule: SigmaRule) -> dict[str, Any]: + logsource = rule.logsource + return { + "rule_id": rule.id, + "title": rule.title, + "status": rule.status, + "level": rule.level, + "tags": list(rule.tags), + "product": logsource.get("product"), + "category": logsource.get("category"), + "service": logsource.get("service"), + "definition": logsource.get("definition"), + "logsource": dict(logsource), + "description": rule.description, + "author": rule.author, + "date": rule.date, + "modified": rule.modified, + "falsepositives": list(rule.falsepositives), + "references": list(rule.references), + } + + +def rich_rule_chunks(rule: SigmaRule) -> list[Chunk]: + metadata = rule_metadata(rule) + source = rule.file_path or rule.id + chunks: list[Chunk] = [] + + def add(chunk_type: str, text: str) -> None: + text = text.strip() + if text: + chunks.append( + Chunk( + text=text, + source_file=source, + chunk_type=chunk_type, + chunk_index=len(chunks), + metadata=metadata, + ) + ) + + logsource = format_logsource(rule.logsource) if rule.logsource else "" + executive = [f"Title: {rule.title}", f"Rule ID: {rule.id}"] + if rule.status: + executive.append(f"Status: {rule.status}") + if rule.level: + executive.append(f"Level: {rule.level}") + if rule.description: + executive.append(f"Description: {rule.description}") + if logsource: + executive.append(f"Logsource: {logsource}") + if rule.condition: + executive.append(f"Condition: {rule.condition}") + if rule.tags: + executive.append(f"Tags: {', '.join(rule.tags)}") + add("executive_summary", "\n".join(executive)) + + lifecycle = [f"Status: {rule.status or 'unknown'}"] + if rule.author: + lifecycle.append(f"Author: {rule.author}") + if rule.date: + lifecycle.append(f"Date: {rule.date}") + if rule.modified: + lifecycle.append(f"Modified: {rule.modified}") + if rule.license: + lifecycle.append(f"License: {rule.license}") + if rule.references: + lifecycle.append(f"References: {', '.join(rule.references)}") + add("metadata_lifecycle", "\n".join(lifecycle)) + + if rule.logsource: + context = [f"Logsource: {logsource}"] + for key in ("product", "category", "service", "definition"): + value = rule.logsource.get(key) + if value: + context.append(f"{key.capitalize()}: {value}") + add("logsource_context", "\n".join(context)) + + condition_lines: list[str] = [] + if rule.condition: + condition_lines.append(f"Condition: {rule.condition}") + detection_keys = [str(key) for key in rule.detection if key != "condition"] + if detection_keys: + condition_lines.append(f"Detection blocks: {', '.join(detection_keys)}") + if condition_lines: + add("condition", "\n".join(condition_lines)) + + for name, value in rule.detection.items(): + if name == "condition": + continue + rendered = _format_value(value) + if not rendered: + continue + add( + f"detection_block:{name}", + f"Rule: {rule.title}\nDetection block: {name}\n{rendered}", + ) + + if rule.falsepositives: + add("false_positives", "False positives: " + "; ".join(rule.falsepositives)) + return chunks + + +def chunk_rule(rule: SigmaRule, mode: str = "flat") -> list[Chunk]: + if mode == "flat": + text = flat_rule_text(rule).strip() + if not text: + return [] + return [ + Chunk( + text=text, + source_file=rule.file_path or rule.id, + chunk_type="rule", + chunk_index=0, + metadata=rule_metadata(rule), + ) + ] + if mode == "rich": + return rich_rule_chunks(rule) + raise ValueError("mode must be 'flat' or 'rich'") + + +def format_logsource(logsource: Mapping[str, Any]) -> str: + ordered = ("product", "category", "service", "definition") + parts = [f"{key}={logsource[key]}" for key in ordered if logsource.get(key)] + if not parts: + parts = [f"{key}={value}" for key, value in logsource.items() if value not in (None, "")] + return " ".join(str(part) for part in parts) or "unknown" + + +def _as_mapping(value: object) -> dict[str, Any]: + return dict(value) if isinstance(value, Mapping) else {} + + +def _as_str_list(value: object) -> list[str]: + if value is None: + return [] + if isinstance(value, (list, tuple, set)): + return [str(item) for item in value] + return [str(value)] + + +def _as_optional_str(value: object) -> str | None: + if value is None: + return None + text = str(value).strip() + return text or None + + +def _as_optional_int(value: object) -> int | None: + if value is None: + return None + if isinstance(value, bool): + return int(value) + if isinstance(value, int): + return value + if isinstance(value, float): + return int(value) + if isinstance(value, str): + try: + return int(value.strip()) + except ValueError: + return None + return None + + +def _normalize_date(value: object) -> str | None: + if value is None: + return None + if isinstance(value, (date, datetime)): + return value.isoformat() + text = str(value).strip() + return text or None + + +def _format_value(value: Any) -> str: + if value is None: + return "" + if isinstance(value, (list, tuple, set)): + return ", ".join(_format_value(item) for item in value) + if isinstance(value, Mapping): + return " ".join( + f"{key} {_format_value(item)}".strip() for key, item in value.items() + ) + return str(value) diff --git a/sigmaforge/ingestion/text.py b/sigmaforge/ingestion/text.py new file mode 100644 index 0000000..d146671 --- /dev/null +++ b/sigmaforge/ingestion/text.py @@ -0,0 +1,33 @@ +from __future__ import annotations + +import re + +_WHITESPACE = re.compile(r"\s+") + + +def split_text(text: str, chunk_size: int = 1024, chunk_overlap: int = 100) -> list[str]: + if chunk_size <= 0: + raise ValueError("chunk_size must be positive") + if chunk_overlap < 0 or chunk_overlap >= chunk_size: + raise ValueError("chunk_overlap must be non-negative and smaller than chunk_size") + normalized = _WHITESPACE.sub(" ", text).strip() + if not normalized: + return [] + if len(normalized) <= chunk_size: + return [normalized] + + chunks: list[str] = [] + start = 0 + while start < len(normalized): + end = min(start + chunk_size, len(normalized)) + if end < len(normalized): + boundary = normalized.rfind(" ", start + chunk_size // 2, end) + if boundary != -1: + end = boundary + chunk = normalized[start:end].strip() + if chunk: + chunks.append(chunk) + if end >= len(normalized): + break + start = max(end - chunk_overlap, start + 1) + return chunks diff --git a/sigmaforge/llm/__init__.py b/sigmaforge/llm/__init__.py new file mode 100644 index 0000000..b619c3a --- /dev/null +++ b/sigmaforge/llm/__init__.py @@ -0,0 +1,3 @@ +from sigmaforge.llm.client import LlamaClient + +__all__ = ["LlamaClient"] diff --git a/sigmaforge/llm/client.py b/sigmaforge/llm/client.py new file mode 100644 index 0000000..aeed36c --- /dev/null +++ b/sigmaforge/llm/client.py @@ -0,0 +1,129 @@ +from __future__ import annotations + +import json +from collections.abc import AsyncIterator, Sequence +from typing import Any, Mapping + +import httpx + +from sigmaforge.errors import LlamaError + + +class LlamaClient: + def __init__( + self, + base_url: str, + model: str = "sigma", + api_key: str = "sigma-key", + timeout: float = 120.0, + transport: httpx.AsyncBaseTransport | None = None, + ) -> None: + self.base_url = base_url.rstrip("/") + self.model = model + self.api_key = api_key + self.timeout = timeout + self.transport = transport + + def _headers(self) -> dict[str, str]: + return { + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json", + } + + async def _post_json(self, path: str, payload: dict[str, Any]) -> dict[str, Any]: + url = f"{self.base_url}{path}" + try: + async with httpx.AsyncClient(transport=self.transport, timeout=self.timeout) as client: + response = await client.post(url, json=payload, headers=self._headers()) + response.raise_for_status() + data: dict[str, Any] = response.json() + return data + except httpx.HTTPError as exc: + raise LlamaError(str(exc)) from exc + except (json.JSONDecodeError, KeyError, TypeError) as exc: + raise LlamaError(str(exc)) from exc + + async def complete( + self, + prompt: str, + *, + temperature: float = 0.3, + max_tokens: int = 512, + stop: str | Sequence[str] | None = None, + ) -> str: + payload: dict[str, Any] = { + "model": self.model, + "prompt": prompt, + "temperature": temperature, + "max_tokens": max_tokens, + } + if stop is not None: + payload["stop"] = stop + data = await self._post_json("/v1/completions", payload) + choices = data.get("choices") or [] + if not choices: + raise LlamaError("completion response contains no choices") + return str(choices[0].get("text", "")) + + async def chat( + self, + messages: Sequence[Mapping[str, Any]], + *, + temperature: float = 0.3, + max_tokens: int = 512, + stop: str | Sequence[str] | None = None, + ) -> str: + payload: dict[str, Any] = { + "model": self.model, + "messages": list(messages), + "temperature": temperature, + "max_tokens": max_tokens, + } + if stop is not None: + payload["stop"] = stop + data = await self._post_json("/v1/chat/completions", payload) + choices = data.get("choices") or [] + if not choices: + raise LlamaError("chat response contains no choices") + return str(choices[0].get("message", {}).get("content", "")) + + async def stream_chat( + self, + messages: Sequence[Mapping[str, Any]], + *, + temperature: float = 0.3, + max_tokens: int = 512, + stop: str | Sequence[str] | None = None, + ) -> AsyncIterator[str]: + payload: dict[str, Any] = { + "model": self.model, + "messages": list(messages), + "temperature": temperature, + "max_tokens": max_tokens, + "stream": True, + } + if stop is not None: + payload["stop"] = stop + url = f"{self.base_url}/v1/chat/completions" + try: + async with httpx.AsyncClient(transport=self.transport, timeout=self.timeout) as client: + async with client.stream("POST", url, json=payload, headers=self._headers()) as response: + response.raise_for_status() + async for line in response.aiter_lines(): + if not line.startswith("data:"): + continue + data = line[5:].strip() + if data == "[DONE]": + break + try: + event = json.loads(data) + except json.JSONDecodeError: + continue + choices = event.get("choices") or [] + if not choices: + continue + chunk = choices[0].get("delta", {}).get("content") + if chunk: + yield str(chunk) + except httpx.HTTPError as exc: + raise LlamaError(str(exc)) from exc diff --git a/sigmaforge/qdrant/__init__.py b/sigmaforge/qdrant/__init__.py new file mode 100644 index 0000000..4522233 --- /dev/null +++ b/sigmaforge/qdrant/__init__.py @@ -0,0 +1,18 @@ +from sigmaforge.qdrant.collections import ( + DEFAULT_COLLECTIONS, + DEFAULT_VECTOR_SIZE, + QdrantCollectionManager, + SPARSE_NAME, +) +from sigmaforge.qdrant.connection import QdrantConnection +from sigmaforge.qdrant.store import QdrantStore, SearchHit + +__all__ = [ + "DEFAULT_COLLECTIONS", + "DEFAULT_VECTOR_SIZE", + "QdrantCollectionManager", + "QdrantConnection", + "QdrantStore", + "SearchHit", + "SPARSE_NAME", +] diff --git a/sigmaforge/qdrant/collections.py b/sigmaforge/qdrant/collections.py new file mode 100644 index 0000000..65cf757 --- /dev/null +++ b/sigmaforge/qdrant/collections.py @@ -0,0 +1,144 @@ +from __future__ import annotations + +from qdrant_client import models + +from sigmaforge.errors import QdrantError +from sigmaforge.qdrant.connection import QdrantConnection + +DEFAULT_COLLECTIONS: tuple[str, ...] = ("sigma_rules", "sigma_docs", "sigma_spec") +DEFAULT_VECTOR_SIZE = 384 +SPARSE_NAME = "text-sparse" + +ON_DISK_BY_COLLECTION: dict[str, bool] = { + "sigma_rules": False, + "sigma_docs": True, + "sigma_spec": True, +} + +HNSW_BY_COLLECTION: dict[str, models.HnswConfigDiff] = { + "sigma_rules": models.HnswConfigDiff( + m=16, + ef_construct=200, + full_scan_threshold=10_000, + ), + "sigma_docs": models.HnswConfigDiff( + m=16, + ef_construct=100, + full_scan_threshold=10_000, + ), + "sigma_spec": models.HnswConfigDiff( + m=16, + ef_construct=100, + full_scan_threshold=10_000, + ), +} + +PAYLOAD_INDEXES: dict[str, models.PayloadSchemaType] = { + "source": models.PayloadSchemaType.KEYWORD, + "source_file": models.PayloadSchemaType.KEYWORD, + "collection": models.PayloadSchemaType.KEYWORD, + "chunk_type": models.PayloadSchemaType.KEYWORD, + "rule_id": models.PayloadSchemaType.KEYWORD, + "title": models.PayloadSchemaType.KEYWORD, + "author": models.PayloadSchemaType.KEYWORD, + "level": models.PayloadSchemaType.KEYWORD, + "status": models.PayloadSchemaType.KEYWORD, + "product": models.PayloadSchemaType.KEYWORD, + "category": models.PayloadSchemaType.KEYWORD, + "service": models.PayloadSchemaType.KEYWORD, + "modified": models.PayloadSchemaType.KEYWORD, + "tags": models.PayloadSchemaType.KEYWORD, + "references": models.PayloadSchemaType.KEYWORD, +} + + +class QdrantCollectionManager: + def __init__( + self, + connection: QdrantConnection, + vector_size: int = DEFAULT_VECTOR_SIZE, + enable_hybrid: bool = True, + collections: tuple[str, ...] = DEFAULT_COLLECTIONS, + sparse_name: str = SPARSE_NAME, + ) -> None: + self.connection = connection + self.vector_size = vector_size + self.enable_hybrid = enable_hybrid + self.collections = collections + self.sparse_name = sparse_name + + def client(self): + return self.connection.client() + + def exists(self, name: str) -> bool: + try: + collections = self.client().get_collections().collections + except Exception as exc: + raise QdrantError(str(exc)) from exc + return any(collection.name == name for collection in collections) + + def ensure(self, names: tuple[str, ...] | None = None) -> None: + for name in names or self.collections: + self.ensure_collection(name) + + def ensure_collection(self, name: str) -> None: + if self.exists(name): + return + vectors_config = {"dense": self._dense_params(name)} + sparse_vectors_config = None + if self.enable_hybrid: + sparse_vectors_config = { + self.sparse_name: models.SparseVectorParams( + modifier=models.Modifier.IDF, + ) + } + try: + self.client().create_collection( + collection_name=name, + vectors_config=vectors_config, + sparse_vectors_config=sparse_vectors_config, + ) + except Exception as exc: + raise QdrantError(str(exc)) from exc + for field_name, field_schema in PAYLOAD_INDEXES.items(): + try: + self.client().create_payload_index( + collection_name=name, + field_name=field_name, + field_schema=field_schema, + ) + except Exception as exc: + raise QdrantError(str(exc)) from exc + + def delete_collection(self, name: str) -> None: + if not self.exists(name): + return + try: + self.client().delete_collection(collection_name=name) + except Exception as exc: + raise QdrantError(str(exc)) from exc + + def count(self, name: str) -> int: + try: + return int(self.client().count(collection_name=name, exact=True).count) + except Exception as exc: + raise QdrantError(str(exc)) from exc + + def collection_info(self, name: str) -> dict[str, object]: + try: + info = self.client().get_collection(collection_name=name) + except Exception as exc: + raise QdrantError(str(exc)) from exc + return { + "name": name, + "points_count": info.points_count, + "status": str(info.status), + } + + def _dense_params(self, name: str) -> models.VectorParams: + return models.VectorParams( + size=self.vector_size, + distance=models.Distance.COSINE, + hnsw_config=HNSW_BY_COLLECTION.get(name), + on_disk=ON_DISK_BY_COLLECTION.get(name, False), + ) diff --git a/sigmaforge/qdrant/connection.py b/sigmaforge/qdrant/connection.py new file mode 100644 index 0000000..d2abf4d --- /dev/null +++ b/sigmaforge/qdrant/connection.py @@ -0,0 +1,42 @@ +from __future__ import annotations + +from qdrant_client import QdrantClient + + +class QdrantConnection: + def __init__( + self, + host: str = "127.0.0.1", + port: int = 6333, + timeout: float = 30.0, + location: str | None = None, + ) -> None: + self.host = host + self.port = port + self.timeout = timeout + self.location = location + self._client: QdrantClient | None = None + + def client(self) -> QdrantClient: + if self._client is None: + if self.location is not None: + self._client = QdrantClient(location=self.location, timeout=int(self.timeout)) + else: + self._client = QdrantClient( + host=self.host, + port=self.port, + timeout=int(self.timeout), + ) + return self._client + + def health(self) -> bool: + try: + self.client().get_collections() + except Exception: + return False + return True + + def close(self) -> None: + if self._client is not None: + self._client.close() + self._client = None diff --git a/sigmaforge/qdrant/store.py b/sigmaforge/qdrant/store.py new file mode 100644 index 0000000..3046031 --- /dev/null +++ b/sigmaforge/qdrant/store.py @@ -0,0 +1,166 @@ +from __future__ import annotations + +import uuid +from dataclasses import dataclass +from typing import Any, Mapping, Sequence + +from qdrant_client import models + +from sigmaforge.errors import QdrantError +from sigmaforge.qdrant.collections import QdrantCollectionManager + + +@dataclass(frozen=True, slots=True) +class SearchHit: + id: str + score: float + payload: dict[str, Any] + + +class QdrantStore: + def __init__( + self, + manager: QdrantCollectionManager, + sparse_name: str | None = None, + ) -> None: + self.manager = manager + self.sparse_name = sparse_name or manager.sparse_name + + @staticmethod + def make_point_id(collection: str, payload: Mapping[str, Any]) -> str: + source_file = payload.get("source_file") or payload.get("source") + chunk_type = payload.get("chunk_type") + if source_file and chunk_type: + return str(uuid.uuid5(uuid.NAMESPACE_URL, f"{collection}:{source_file}:{chunk_type}")) + return str(uuid.uuid4()) + + def upsert_texts( + self, + collection: str, + texts: Sequence[str], + dense_vectors: Sequence[Sequence[float]], + payloads: Sequence[Mapping[str, Any]], + sparse_vectors: Sequence[tuple[Sequence[int], Sequence[float]]] | None = None, + ids: Sequence[str] | None = None, + ) -> None: + if len(texts) != len(dense_vectors) or len(texts) != len(payloads): + raise ValueError("texts, dense_vectors, and payloads must have the same length") + points: list[models.PointStruct] = [] + for index, text in enumerate(texts): + payload = dict(payloads[index]) + payload["text"] = text + point_id = str(ids[index]) if ids is not None else self.make_point_id(collection, payload) + vector: dict[str, Any] = {"dense": list(dense_vectors[index])} + if self.manager.enable_hybrid: + indices: Sequence[int] = [] + values: Sequence[float] = [] + if sparse_vectors is not None: + indices, values = sparse_vectors[index] + vector[self.sparse_name] = models.SparseVector( + indices=[int(item) for item in indices], + values=[float(item) for item in values], + ) + points.append( + models.PointStruct( + id=point_id, + vector=vector, + payload=payload, + ) + ) + try: + self.manager.client().upsert( + collection_name=collection, + points=points, + wait=True, + ) + except Exception as exc: + raise QdrantError(str(exc)) from exc + + def replace_source_points( + self, + collection: str, + source_file: str, + texts: Sequence[str], + dense_vectors: Sequence[Sequence[float]], + payloads: Sequence[Mapping[str, Any]], + sparse_vectors: Sequence[tuple[Sequence[int], Sequence[float]]] | None = None, + ids: Sequence[str] | None = None, + ) -> None: + self.delete_by_source(collection, source_file) + self.upsert_texts( + collection=collection, + texts=texts, + dense_vectors=dense_vectors, + payloads=payloads, + sparse_vectors=sparse_vectors, + ids=ids, + ) + + def delete_by_source(self, collection: str, source_file: str) -> None: + filter = models.Filter( + must=[ + models.FieldCondition( + key="source_file", + match=models.MatchValue(value=source_file), + ) + ] + ) + try: + self.manager.client().delete( + collection_name=collection, + points_selector=filter, + wait=True, + ) + except Exception as exc: + raise QdrantError(str(exc)) from exc + + def query_dense( + self, + collection: str, + vector: Sequence[float], + limit: int = 10, + query_filter: models.Filter | None = None, + ) -> list[SearchHit]: + try: + response = self.manager.client().query_points( + collection_name=collection, + query=list(vector), + using="dense", + limit=limit, + query_filter=query_filter, + ) + except Exception as exc: + raise QdrantError(str(exc)) from exc + return [self._to_hit(point) for point in response.points] + + def query_sparse( + self, + collection: str, + indices: Sequence[int], + values: Sequence[float], + limit: int = 10, + query_filter: models.Filter | None = None, + ) -> list[SearchHit]: + if not self.manager.enable_hybrid or not indices: + return [] + try: + response = self.manager.client().query_points( + collection_name=collection, + query=models.SparseVector( + indices=[int(item) for item in indices], + values=[float(item) for item in values], + ), + using=self.sparse_name, + limit=limit, + query_filter=query_filter, + ) + except Exception as exc: + raise QdrantError(str(exc)) from exc + return [self._to_hit(point) for point in response.points] + + def _to_hit(self, point: models.ScoredPoint) -> SearchHit: + return SearchHit( + id=str(point.id), + score=float(point.score), + payload=dict(point.payload or {}), + ) diff --git a/sigmaforge/rag/__init__.py b/sigmaforge/rag/__init__.py new file mode 100644 index 0000000..c484b49 --- /dev/null +++ b/sigmaforge/rag/__init__.py @@ -0,0 +1,11 @@ +from __future__ import annotations + +from sigmaforge.rag.models import RAGAnswer +from sigmaforge.rag.pipeline import RAGPipeline +from sigmaforge.rag.prompts import render_search_prompt + +__all__ = [ + "RAGAnswer", + "RAGPipeline", + "render_search_prompt", +] diff --git a/sigmaforge/rag/models.py b/sigmaforge/rag/models.py new file mode 100644 index 0000000..5e4093a --- /dev/null +++ b/sigmaforge/rag/models.py @@ -0,0 +1,24 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any + +from sigmaforge.search import SearchResult + + +@dataclass(frozen=True, slots=True) +class RAGAnswer: + query: str + answer: str + sources: tuple[SearchResult, ...] + model: str | None = None + fallback: bool = False + + def to_dict(self) -> dict[str, Any]: + return { + "query": self.query, + "answer": self.answer, + "sources": [source.to_dict() for source in self.sources], + "model": self.model, + "fallback": self.fallback, + } diff --git a/sigmaforge/rag/pipeline.py b/sigmaforge/rag/pipeline.py new file mode 100644 index 0000000..ae8809d --- /dev/null +++ b/sigmaforge/rag/pipeline.py @@ -0,0 +1,185 @@ +from __future__ import annotations + +import logging +from collections.abc import AsyncIterator, Mapping, Sequence +from typing import Any, Protocol + +from sigmaforge.errors import LlamaError +from sigmaforge.rag.models import RAGAnswer +from sigmaforge.rag.prompts import render_search_prompt +from sigmaforge.search import SearchResult, format_context + +logger = logging.getLogger(__name__) + + +class SearchEngineLike(Protocol): + def search( + self, + query: str, + *, + top_k: int | None = None, + collections: Sequence[str] | None = None, + filters: Mapping[str, str] | None = None, + extra_filter: Any | None = None, + ) -> list[SearchResult]: + ... + + +class LlamaClientLike(Protocol): + model: str + + async def chat( + self, + messages: Sequence[Mapping[str, Any]], + *, + temperature: float = 0.3, + max_tokens: int = 512, + stop: str | Sequence[str] | None = None, + ) -> str: + ... + + def stream_chat( + self, + messages: Sequence[Mapping[str, Any]], + *, + temperature: float = 0.3, + max_tokens: int = 512, + stop: str | Sequence[str] | None = None, + ) -> AsyncIterator[str]: + ... + + +class RAGPipeline: + def __init__( + self, + search_engine: SearchEngineLike, + llm_client: LlamaClientLike, + *, + top_k: int = 5, + context_max_results: int = 5, + context_max_chars: int = 1200, + temperature: float = 0.3, + max_tokens: int = 512, + ) -> None: + self._search_engine = search_engine + self._llm_client = llm_client + self._top_k = top_k + self._context_max_results = context_max_results + self._context_max_chars = context_max_chars + self._temperature = temperature + self._max_tokens = max_tokens + + async def answer( + self, + query: str, + *, + top_k: int | None = None, + collections: Sequence[str] | None = None, + filters: Mapping[str, str] | None = None, + extra_filter: Any | None = None, + ) -> RAGAnswer: + limit = top_k if top_k is not None else self._top_k + results = self._search_engine.search( + query, + top_k=limit, + collections=collections, + filters=filters, + extra_filter=extra_filter, + ) + if not results: + return RAGAnswer( + query=query, + answer=self._fallback_answer(results), + sources=(), + model=self._llm_client.model, + fallback=True, + ) + + context = format_context( + results, + max_results=self._context_max_results, + max_chars=self._context_max_chars, + ) + prompt = render_search_prompt(context, query) + messages: list[dict[str, str]] = [ + {"role": "system", "content": prompt}, + {"role": "user", "content": query}, + ] + + try: + response = await self._llm_client.chat( + messages, + temperature=self._temperature, + max_tokens=self._max_tokens, + ) + except LlamaError as exc: + logger.warning("RAG LLM generation failed: %s", exc) + return RAGAnswer( + query=query, + answer=self._fallback_answer(results), + sources=tuple(results), + model=self._llm_client.model, + fallback=True, + ) + + return RAGAnswer( + query=query, + answer=response.strip(), + sources=tuple(results), + model=self._llm_client.model, + fallback=False, + ) + + async def answer_stream( + self, + query: str, + *, + top_k: int | None = None, + collections: Sequence[str] | None = None, + filters: Mapping[str, str] | None = None, + extra_filter: Any | None = None, + ) -> AsyncIterator[str]: + limit = top_k if top_k is not None else self._top_k + results = self._search_engine.search( + query, + top_k=limit, + collections=collections, + filters=filters, + extra_filter=extra_filter, + ) + if not results: + yield self._fallback_answer(results) + return + + context = format_context( + results, + max_results=self._context_max_results, + max_chars=self._context_max_chars, + ) + prompt = render_search_prompt(context, query) + messages: list[dict[str, str]] = [ + {"role": "system", "content": prompt}, + {"role": "user", "content": query}, + ] + + started = False + try: + async for chunk in self._llm_client.stream_chat( + messages, + temperature=self._temperature, + max_tokens=self._max_tokens, + ): + started = True + yield chunk + except LlamaError as exc: + logger.warning("RAG LLM stream failed: %s", exc) + if not started: + yield self._fallback_answer(results) + + def _fallback_answer(self, results: Sequence[SearchResult]) -> str: + if not results: + return "No matching Sigma rules or documents were found." + lines = [ + f"{index}. {result.text[:200]}" for index, result in enumerate(results[:2], start=1) + ] + return "LLM unavailable. Top results:\n" + "\n".join(lines) diff --git a/sigmaforge/rag/prompts.py b/sigmaforge/rag/prompts.py new file mode 100644 index 0000000..9906db9 --- /dev/null +++ b/sigmaforge/rag/prompts.py @@ -0,0 +1,26 @@ +from __future__ import annotations + +SEARCH_PROMPT = """\ +You are a cybersecurity expert helping SOC analysts with detection questions and Sigma specification lookups. + +Search Results (from vector search over Sigma rules, documentation, and specification docs): +{search_results} + +Question: {question} + +Task: Answer the user's question using ONLY the search results above. The results are ordered by relevance. Focus primarily on the first result. If later results discuss a different topic, ignore them. Cite specific rule names, detection logic, specification attributes, and file paths. If the search results do not contain enough information, say so clearly — do NOT guess or use outside knowledge. + +When results include Sigma specification content, mention: +- The exact Sigma attribute or field name, its purpose, required/optional status, valid values, and concrete YAML examples from the spec. + +When results include Sigma rules, mention: +- Rule names and detection logic +- MITRE ATT&CK mapping when available +- False positive considerations + +Format your answer clearly with Markdown. Keep it concise and scannable. +""" + + +def render_search_prompt(search_results: str, question: str) -> str: + return SEARCH_PROMPT.format(search_results=search_results, question=question) diff --git a/sigmaforge/search/__init__.py b/sigmaforge/search/__init__.py new file mode 100644 index 0000000..a24a0e4 --- /dev/null +++ b/sigmaforge/search/__init__.py @@ -0,0 +1,38 @@ +from sigmaforge.search.context import format_context, format_result, get_citation +from sigmaforge.search.engine import ( + ALPHA_BY_COLLECTION, + DEFAULT_ALPHA, + DEFAULT_TOP_K, + SIMILARITY_THRESHOLD, + DenseEncoderLike, + SearchEngine, + SparseEncoderLike, +) +from sigmaforge.search.filters import ( + FILTER_KEYS, + LIST_FILTER_KEYS, + build_qdrant_filter, + parse_query_filters, +) +from sigmaforge.search.fusion import RRF_K_DEFAULT, reciprocal_rank_fusion +from sigmaforge.search.models import SearchResult + +__all__ = [ + "ALPHA_BY_COLLECTION", + "DEFAULT_ALPHA", + "DEFAULT_TOP_K", + "FILTER_KEYS", + "LIST_FILTER_KEYS", + "RRF_K_DEFAULT", + "SIMILARITY_THRESHOLD", + "DenseEncoderLike", + "SearchEngine", + "SearchResult", + "SparseEncoderLike", + "build_qdrant_filter", + "format_context", + "format_result", + "get_citation", + "parse_query_filters", + "reciprocal_rank_fusion", +] diff --git a/sigmaforge/search/context.py b/sigmaforge/search/context.py new file mode 100644 index 0000000..c840eee --- /dev/null +++ b/sigmaforge/search/context.py @@ -0,0 +1,83 @@ +from __future__ import annotations + +from collections.abc import Sequence +from typing import Any + +from sigmaforge.search.models import SearchResult + + +def format_result(result: SearchResult) -> dict[str, Any]: + metadata = result.metadata + base: dict[str, Any] = { + "id": result.id, + "text": result.text, + "score": result.score, + "collection": result.collection, + "source_file": metadata.get("source_file", ""), + "file_path": metadata.get("file_path", metadata.get("source_file", "")), + "line_start": metadata.get("line_start", metadata.get("line_number", "")), + "metadata": metadata, + } + + if result.collection == "sigma_rules": + base.update( + { + "rule_id": metadata.get("rule_id", ""), + "title": metadata.get("title", ""), + "level": metadata.get("level", ""), + "status": metadata.get("status", ""), + "chunk_type": metadata.get("chunk_type", ""), + "product": metadata.get("product", ""), + "category": metadata.get("category", ""), + } + ) + elif result.collection == "sigma_docs": + base.update( + { + "doc_type": metadata.get("doc_type", ""), + "heading_text": metadata.get("heading_text", ""), + "heading_level": metadata.get("heading_level", 0), + "original_url": metadata.get("original_url", ""), + "source_rule_id": metadata.get("rule_id", ""), + } + ) + + return base + + +def get_citation(result: SearchResult) -> str: + metadata = result.metadata + source = metadata.get("file_path") or metadata.get("source_file") or "" + line = metadata.get("line_start") or metadata.get("line_number") or "" + if source and line: + return f"{source}:{line}" + return str(source) + + +def format_context( + results: Sequence[SearchResult], + *, + max_results: int = 5, + max_chars: int = 1200, +) -> str: + blocks: list[str] = [] + budget = max_chars + for index, result in enumerate(results[: max(0, max_results)], start=1): + metadata = result.metadata + source = metadata.get("source_file") or metadata.get("file_path") or result.collection + title = metadata.get("title") + header = f"[{index}] {result.collection} / {source}" + if title: + header += f" / {title}" + available = budget - len(header) - 1 + if available <= 0: + break + text = result.text + if len(text) > available: + if available <= 3: + break + text = text[: available - 3].rstrip() + "..." + block = f"{header}\n{text}" + blocks.append(block) + budget -= len(block) + return "\n\n".join(blocks) diff --git a/sigmaforge/search/engine.py b/sigmaforge/search/engine.py new file mode 100644 index 0000000..04a0d94 --- /dev/null +++ b/sigmaforge/search/engine.py @@ -0,0 +1,261 @@ +from __future__ import annotations + +from collections.abc import Mapping, Sequence +from typing import Any, Protocol + +from qdrant_client import models + +from sigmaforge.embed.sparse import SparseVector +from sigmaforge.errors import SearchError +from sigmaforge.qdrant import QdrantStore, SearchHit +from sigmaforge.search.context import format_context, format_result, get_citation +from sigmaforge.search.filters import build_qdrant_filter, parse_query_filters +from sigmaforge.search.fusion import RRF_K_DEFAULT, reciprocal_rank_fusion +from sigmaforge.search.models import SearchResult + +DEFAULT_TOP_K = 15 +SIMILARITY_THRESHOLD = 0.0 +DEFAULT_ALPHA = 0.5 + +ALPHA_BY_COLLECTION: dict[str, float] = { + "sigma_rules": 0.5, + "sigma_docs": 0.7, + "sigma_spec": 0.3, +} + + +class DenseEncoderLike(Protocol): + def encode(self, texts: Sequence[str], *, is_query: bool = False) -> list[list[float]]: + ... + + +class SparseEncoderLike(Protocol): + def encode_text(self, text: str) -> SparseVector: + ... + + +class SearchEngine: + def __init__( + self, + store: QdrantStore, + dense_encoder: DenseEncoderLike, + sparse_encoder: SparseEncoderLike | None = None, + *, + collections: Sequence[str] | None = None, + top_k: int = DEFAULT_TOP_K, + similarity_threshold: float = SIMILARITY_THRESHOLD, + alpha: float = DEFAULT_ALPHA, + alpha_by_collection: Mapping[str, float] | None = None, + rrf_k: int = RRF_K_DEFAULT, + rrf_weights: Mapping[str, float] | None = None, + ) -> None: + self._store = store + self._dense_encoder = dense_encoder + self._sparse_encoder = sparse_encoder + self._collections = tuple(collections or store.manager.collections) + self._top_k = top_k + self._similarity_threshold = similarity_threshold + self._alpha = alpha + self._alpha_by_collection = dict(alpha_by_collection or {}) + self._rrf_k = rrf_k + self._rrf_weights = dict(rrf_weights or {}) + + @property + def collections(self) -> tuple[str, ...]: + return self._collections + + def search( + self, + query: str, + *, + top_k: int | None = None, + collections: Sequence[str] | None = None, + filters: Mapping[str, str] | None = None, + extra_filter: models.Filter | None = None, + ) -> list[SearchResult]: + normalized = self._normalize_query(query) + if not normalized: + return [] + + limit = top_k if top_k is not None else self._top_k + if limit <= 0: + return [] + + resolved_filters = dict(filters or {}) + inline_filters, clean_query = parse_query_filters(normalized) + resolved_filters.update(inline_filters) + embed_query = clean_query if clean_query else normalized + + selected = list(collections or self._collections) + if ( + "references" in resolved_filters + and "sigma_docs" in self._collections + and "sigma_docs" not in selected + ): + selected.append("sigma_docs") + if not selected: + return [] + + qdrant_filter = build_qdrant_filter(resolved_filters, extra_filter) + dense_vector = self._dense_vector(embed_query) + sparse_vector = self._sparse_vector(embed_query) + + per_collection_limit = max(limit * 2, 10) + rankings: dict[str, list[SearchResult]] = {} + for collection in selected: + results = self._search_collection( + collection, + embed_query, + per_collection_limit, + qdrant_filter, + dense_vector, + sparse_vector, + ) + if results: + rankings[collection] = results + if not rankings: + return [] + return reciprocal_rank_fusion( + rankings, + k=self._rrf_k, + weights=self._rrf_weights, + limit=limit, + ) + + def search_collection( + self, + collection: str, + query: str, + *, + top_k: int | None = None, + filters: Mapping[str, str] | None = None, + extra_filter: models.Filter | None = None, + ) -> list[SearchResult]: + normalized = self._normalize_query(query) + if not normalized: + return [] + + limit = top_k if top_k is not None else self._top_k + if limit <= 0: + return [] + + resolved_filters = dict(filters or {}) + inline_filters, clean_query = parse_query_filters(normalized) + resolved_filters.update(inline_filters) + embed_query = clean_query if clean_query else normalized + + qdrant_filter = build_qdrant_filter(resolved_filters, extra_filter) + return self._search_collection( + collection, + embed_query, + limit, + qdrant_filter, + self._dense_vector(embed_query), + self._sparse_vector(embed_query), + ) + + def format_result(self, result: SearchResult) -> dict[str, Any]: + return format_result(result) + + def get_citation(self, result: SearchResult) -> str: + return get_citation(result) + + def format_context( + self, + results: Sequence[SearchResult], + *, + max_results: int = 5, + max_chars: int = 1200, + ) -> str: + return format_context(results, max_results=max_results, max_chars=max_chars) + + def _search_collection( + self, + collection: str, + query: str, + limit: int, + qdrant_filter: models.Filter | None, + dense_vector: list[float] | None, + sparse_vector: SparseVector | None, + ) -> list[SearchResult]: + alpha = float(self._alpha_by_collection.get(collection, self._alpha)) + dense_weight = min(max(alpha, 0.0), 1.0) + sparse_weight = max(0.0, 1.0 - dense_weight) + if self._sparse_encoder is None: + dense_weight = 1.0 + sparse_weight = 0.0 + + dense_hits: list[SearchHit] = [] + if dense_weight > 0.0 and dense_vector: + dense_hits = self._store.query_dense( + collection, + dense_vector, + limit=limit, + query_filter=qdrant_filter, + ) + + sparse_hits: list[SearchHit] = [] + if sparse_weight > 0.0 and sparse_vector is not None and sparse_vector.indices: + sparse_hits = self._store.query_sparse( + collection, + list(sparse_vector.indices), + list(sparse_vector.values), + limit=limit, + query_filter=qdrant_filter, + ) + + dense_hits = [hit for hit in dense_hits if hit.score >= self._similarity_threshold] + sparse_hits = [hit for hit in sparse_hits if hit.score >= self._similarity_threshold] + + rankings: dict[str, list[SearchResult]] = {} + if dense_hits: + rankings["dense"] = [self._to_result(hit, collection) for hit in dense_hits] + if sparse_hits: + rankings["sparse"] = [self._to_result(hit, collection) for hit in sparse_hits] + if not rankings: + return [] + + return reciprocal_rank_fusion( + rankings, + k=self._rrf_k, + weights={"dense": dense_weight, "sparse": sparse_weight}, + limit=limit, + ) + + @staticmethod + def _normalize_query(query: str) -> str: + return query.replace("`", "").rstrip("?").strip() + + def _dense_vector(self, query: str) -> list[float] | None: + if not query.strip(): + return None + try: + vectors = self._dense_encoder.encode([query], is_query=True) + except Exception as exc: + raise SearchError(f"dense encoding failed: {exc}") from exc + if not vectors or not vectors[0]: + return None + return list(vectors[0]) + + def _sparse_vector(self, query: str) -> SparseVector | None: + if self._sparse_encoder is None or not query.strip(): + return None + try: + vector = self._sparse_encoder.encode_text(query) + except Exception as exc: + raise SearchError(f"sparse encoding failed: {exc}") from exc + if not vector.indices: + return None + return vector + + @staticmethod + def _to_result(hit: SearchHit, collection: str) -> SearchResult: + payload = dict(hit.payload or {}) + resolved_collection = str(payload.get("collection") or collection) + return SearchResult( + id=str(hit.id), + collection=resolved_collection, + text=str(payload.get("text") or ""), + score=float(hit.score or 0.0), + payload=payload, + ) diff --git a/sigmaforge/search/filters.py b/sigmaforge/search/filters.py new file mode 100644 index 0000000..cb8f3b9 --- /dev/null +++ b/sigmaforge/search/filters.py @@ -0,0 +1,84 @@ +from __future__ import annotations + +import re +from collections.abc import Mapping + +from qdrant_client import models + +FILTER_KEYS: frozenset[str] = frozenset( + { + "rule_id", + "title", + "author", + "level", + "status", + "product", + "category", + "service", + "date", + "modified", + "chunk_type", + "collection", + "tags", + "references", + } +) + +LIST_FILTER_KEYS: frozenset[str] = frozenset({"tags", "references"}) + +_FILTER_PATTERN = re.compile(r"(\w+):\s*(\S+)") +_VALUE_CLEAN_RE = re.compile(r"[\s,;:]+") +_WHITESPACE_RE = re.compile(r"\s+") + + +def _clean_value(raw: str) -> str: + value = raw.strip().strip("'\"") + return _VALUE_CLEAN_RE.sub(" ", value).strip() + + +def _remove_known_filter(match: re.Match[str]) -> str: + if match.group(1).lower() in FILTER_KEYS: + return "" + return match.group(0) + + +def parse_query_filters(query: str) -> tuple[dict[str, str], str]: + filters: dict[str, str] = {} + for match in _FILTER_PATTERN.finditer(query): + key = match.group(1).lower() + if key not in FILTER_KEYS: + continue + value = _clean_value(match.group(2)) + if value: + filters[key] = value + cleaned = _FILTER_PATTERN.sub(_remove_known_filter, query) + return filters, _WHITESPACE_RE.sub(" ", cleaned).strip() + + +def _split_values(value: str) -> list[str]: + return [part for part in re.split(r"[,\s]+", value) if part] + + +def build_qdrant_filter( + filters: Mapping[str, str] | None = None, + extra_filter: models.Filter | None = None, +) -> models.Filter | None: + conditions: list[models.Condition] = [] + for key, value in (filters or {}).items(): + if not value: + continue + if key in LIST_FILTER_KEYS: + values = _split_values(value) + if values: + conditions.append( + models.FieldCondition(key=key, match=models.MatchAny(any=values)) + ) + else: + conditions.append( + models.FieldCondition(key=key, match=models.MatchValue(value=value)) + ) + if extra_filter is None: + return models.Filter(must=conditions) if conditions else None + if not conditions: + return extra_filter + return models.Filter(must=[*conditions, extra_filter]) diff --git a/sigmaforge/search/fusion.py b/sigmaforge/search/fusion.py new file mode 100644 index 0000000..dff740a --- /dev/null +++ b/sigmaforge/search/fusion.py @@ -0,0 +1,37 @@ +from __future__ import annotations + +from collections.abc import Mapping, Sequence +from dataclasses import replace + +from sigmaforge.errors import SearchError +from sigmaforge.search.models import SearchResult + +RRF_K_DEFAULT = 60 + + +def reciprocal_rank_fusion( + rankings: Mapping[str, Sequence[SearchResult]], + *, + k: int = RRF_K_DEFAULT, + weights: Mapping[str, float] | None = None, + limit: int | None = None, +) -> list[SearchResult]: + if k <= 0: + raise SearchError("rrf k must be positive") + + resolved_weights = dict(weights or {}) + scores: dict[str, float] = {} + results: dict[str, SearchResult] = {} + for name, ranked_results in rankings.items(): + weight = float(resolved_weights.get(name, 1.0)) + if weight <= 0.0: + continue + for rank, result in enumerate(ranked_results, start=1): + key = result.fusion_key + scores[key] = scores.get(key, 0.0) + weight / (k + rank) + results.setdefault(key, result) + + ranked_keys = sorted(scores, key=lambda key: (-scores[key], key)) + if limit is not None: + ranked_keys = ranked_keys[: max(0, limit)] + return [replace(results[key], score=scores[key]) for key in ranked_keys] diff --git a/sigmaforge/search/models.py b/sigmaforge/search/models.py new file mode 100644 index 0000000..f2bcc25 --- /dev/null +++ b/sigmaforge/search/models.py @@ -0,0 +1,30 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any + + +@dataclass(frozen=True, slots=True) +class SearchResult: + id: str + collection: str + text: str + score: float + payload: dict[str, Any] + + @property + def metadata(self) -> dict[str, Any]: + return {key: value for key, value in self.payload.items() if key != "text"} + + @property + def fusion_key(self) -> str: + return f"{self.collection}:{self.id}" + + def to_dict(self) -> dict[str, Any]: + return { + "id": self.id, + "collection": self.collection, + "text": self.text, + "score": self.score, + "metadata": self.metadata, + } diff --git a/tests/e2e/__init__.py b/tests/e2e/__init__.py new file mode 100644 index 0000000..783d1c8 --- /dev/null +++ b/tests/e2e/__init__.py @@ -0,0 +1 @@ +"""End-to-end UI tests for the legacy FastAPI application.""" diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py new file mode 100644 index 0000000..802d98f --- /dev/null +++ b/tests/e2e/conftest.py @@ -0,0 +1,83 @@ +from __future__ import annotations + +import threading +import time +import urllib.request +from collections.abc import Generator +from typing import Any + +import pytest +import uvicorn +from _pytest.monkeypatch import MonkeyPatch + +from src import main as main_module +from src.infrastructure.database.core import DatabaseServiceCore + +STARTUP_TIMEOUT_SECONDS = 60 + + +def _wait_for_server(base_url: str) -> None: + deadline = time.monotonic() + STARTUP_TIMEOUT_SECONDS + last_error: Exception | None = None + while time.monotonic() < deadline: + try: + with urllib.request.urlopen(f"{base_url}/api/v1/admin/status", timeout=1) as response: + if response.status == 200: + return + except Exception as exc: + last_error = exc + time.sleep(0.2) + message = f"E2E server did not become ready at {base_url}: {last_error}" + raise RuntimeError(message) + + +@pytest.fixture(scope="module") +def e2e_app(tmp_path_factory: pytest.TempPathFactory) -> Generator[Any, None, None]: + sandbox = tmp_path_factory.mktemp("sigmaforge-e2e") + monkey = MonkeyPatch() + monkey.setenv("HF_HUB_OFFLINE", "1") + monkey.setenv("HF_TOKEN", "") + monkey.setenv("TQDM_DISABLE", "1") + monkey.delenv("_SIGMA_SETUP_MODE", raising=False) + monkey.setattr(main_module, "setup_mode", True) + monkey.chdir(sandbox) + + app = main_module.create_app() + config = uvicorn.Config( + app, + host="127.0.0.1", + port=0, + log_level="warning", + lifespan="on", + ) + config.load() + sock = config.bind_socket() + base_url = f"http://127.0.0.1:{sock.getsockname()[1]}" + server = uvicorn.Server(config) + server_thread = threading.Thread(target=server.run, kwargs={"sockets": [sock]}, daemon=True) + server_thread.start() + _wait_for_server(base_url) + app.state.e2e_base_url = base_url + try: + yield app + finally: + server.should_exit = True + server_thread.join(timeout=10) + if sock.fileno() != -1: + sock.close() + monkey.undo() + + +@pytest.fixture(scope="module") +def e2e_base_url(e2e_app: Any) -> str: + return e2e_app.state.e2e_base_url + + +@pytest.fixture(autouse=True) +def e2e_database( + reset_modules: Generator[Any, None, None], e2e_app: Any +) -> Generator[None, None, None]: + db = e2e_app.state.db + DatabaseServiceCore._instance = db + yield + DatabaseServiceCore._instance = db diff --git a/tests/e2e/test_ui_smoke.py b/tests/e2e/test_ui_smoke.py new file mode 100644 index 0000000..7b5ac0f --- /dev/null +++ b/tests/e2e/test_ui_smoke.py @@ -0,0 +1,41 @@ +from __future__ import annotations + +import pytest +from playwright.sync_api import Page + +pytestmark = pytest.mark.e2e + + +def test_chat_page(page: Page, e2e_base_url: str) -> None: + page.goto(f"{e2e_base_url}/chat") + page.wait_for_selector("#message-input", state="visible") + assert page.locator("#chat-welcome").is_visible() + assert page.locator("#send-btn").is_visible() + + +def test_config_page(page: Page, e2e_base_url: str) -> None: + page.goto(f"{e2e_base_url}/config") + page.wait_for_selector("#status-grid", state="visible") + assert page.locator("#status-banner").is_visible() + assert page.locator("#general-section").is_visible() + + +def test_dashboard_page(page: Page, e2e_base_url: str) -> None: + page.goto(f"{e2e_base_url}/dashboard") + page.wait_for_selector("#duckdb-content", state="visible") + + +def test_logs_page(page: Page, e2e_base_url: str) -> None: + page.goto(f"{e2e_base_url}/logs") + page.wait_for_selector("#log-stats", state="visible") + + +def test_local_data_page(page: Page, e2e_base_url: str) -> None: + page.goto(f"{e2e_base_url}/data/local") + page.wait_for_selector("#local-files-body", state="visible") + assert page.locator("#upload-zone").is_visible() + + +def test_setup_redirects_to_config(page: Page, e2e_base_url: str) -> None: + page.goto(f"{e2e_base_url}/setup") + assert page.url.endswith("/config") diff --git a/tests/integration/test_api_smoke.py b/tests/integration/test_api_smoke.py new file mode 100644 index 0000000..d7f4242 --- /dev/null +++ b/tests/integration/test_api_smoke.py @@ -0,0 +1,89 @@ +from __future__ import annotations + +from pathlib import Path + +import pytest +from fastapi.testclient import TestClient + +from src.infrastructure.database import DatabaseService +from src.main import create_app + + +@pytest.fixture +def client(tmp_path: Path): + db = DatabaseService(str(tmp_path / "sigmaforge_smoke.duckdb")) + db.initialize() + app = create_app() + app.state.db = db + yield TestClient(app) + + +def _get_json(client: TestClient, url: str) -> dict: + response = client.get(url) + response.raise_for_status() + return response.json() + + +def test_chat_page(client: TestClient) -> None: + response = client.get("/chat") + assert response.status_code == 200 + assert "SigmaForge" in response.text + + +def test_config_page(client: TestClient) -> None: + response = client.get("/config") + assert response.status_code == 200 + assert "System Status" in response.text + + +def test_dashboard_page(client: TestClient) -> None: + response = client.get("/dashboard") + assert response.status_code == 200 + assert "Database Dashboard" in response.text + + +def test_logs_page(client: TestClient) -> None: + response = client.get("/logs") + assert response.status_code == 200 + assert "System Logs" in response.text + + +def test_local_data_page(client: TestClient) -> None: + response = client.get("/data/local") + assert response.status_code == 200 + assert "Local Files" in response.text + + +def test_setup_redirects_to_config(client: TestClient) -> None: + response = client.get("/setup", follow_redirects=False) + assert response.status_code == 301 + assert response.headers["location"] == "/config" + + +def test_admin_status(client: TestClient) -> None: + payload = _get_json(client, "/api/v1/admin/status") + assert payload["data"]["llama_cpp"]["status"] in {"active", "inactive"} + assert payload["data"]["qdrant"]["status"] in {"active", "inactive"} + + +def test_chat_history(client: TestClient) -> None: + response = client.get("/api/v1/chat/history", headers={"X-Session-ID": "smoke"}) + assert response.status_code == 200 + assert response.json() == [] + + +def test_prompts_list(client: TestClient) -> None: + response = client.get("/api/v1/admin/prompts") + assert response.status_code == 200 + assert isinstance(response.json(), list) + + +def test_dashboard_tables(client: TestClient) -> None: + payload = _get_json(client, "/api/v1/dashboard/tables") + assert "config" in payload["tables"] + + +def test_logging_config(client: TestClient) -> None: + payload = _get_json(client, "/api/v1/config/logging") + assert payload["status"] == "success" + assert "level" in payload["data"] diff --git a/tests/rewrite/test_config.py b/tests/rewrite/test_config.py new file mode 100644 index 0000000..3115c74 --- /dev/null +++ b/tests/rewrite/test_config.py @@ -0,0 +1,59 @@ +from __future__ import annotations + +from pathlib import Path + +from sigmaforge.config import Config, load_config +from sigmaforge.db import ConfigStore, Database + + +def test_config_defaults() -> None: + config = Config() + assert config.qdrant_host == "127.0.0.1" + assert config.qdrant_port == 6333 + assert config.vector_size == 384 + assert config.collections == ("sigma_rules", "sigma_docs", "sigma_spec") + assert config.llama_base_url == "http://127.0.0.1:8080" + assert config.embedding_batch_size == 64 + assert config.hf_offline is True + + +def test_config_apply_env(monkeypatch) -> None: + monkeypatch.setenv("SIGMAFORGE_DATA_DIR", "/tmp/sigmaforge-data") + monkeypatch.setenv("SIGMAFORGE_QDRANT_PORT", "6334") + monkeypatch.setenv("SIGMAFORGE_VECTOR_SIZE", "512") + monkeypatch.setenv("SIGMAFORGE_COLLECTIONS", "sigma_rules,sigma_docs") + monkeypatch.setenv("SIGMAFORGE_EMBEDDING_BATCH_SIZE", "12") + monkeypatch.setenv("HF_HUB_OFFLINE", "0") + + config = Config() + config.apply_env() + + assert config.data_dir == Path("/tmp/sigmaforge-data") + assert config.duckdb_path == Path("/tmp/sigmaforge-data/duckdb/sigmaforge.duckdb") + assert config.qdrant_port == 6334 + assert config.vector_size == 512 + assert config.collections == ("sigma_rules", "sigma_docs") + assert config.embedding_batch_size == 12 + assert config.hf_offline is False + + +def test_load_config_applies_db_overrides() -> None: + db = Database(":memory:") + db.init_schema() + store = ConfigStore(db) + store.set("paths.duckdb", "/tmp/override.duckdb") + store.set("services.qdrant.port", 6335) + store.set("services.llama.base_url", "http://127.0.0.1:9999") + store.set("search.collections", ["sigma_rules", "sigma_spec"]) + store.set("models.embedding.active", "e5-large") + store.set("ingestion.chunk_size", 777) + + config = load_config(env={}, config_store=store) + + assert config.duckdb_path == Path("/tmp/override.duckdb") + assert config.qdrant_port == 6335 + assert config.llama_base_url == "http://127.0.0.1:9999" + assert config.collections == ("sigma_rules", "sigma_spec") + assert config.embedding_model == "e5-large" + assert config.chunk_size == 777 + db.close() diff --git a/tests/rewrite/test_db.py b/tests/rewrite/test_db.py new file mode 100644 index 0000000..509c3fe --- /dev/null +++ b/tests/rewrite/test_db.py @@ -0,0 +1,76 @@ +from __future__ import annotations + +from typing import Iterator + +import pytest + +from sigmaforge.db import ConfigStore, Database, DocRegistry, ModelStore, Task, TaskStore + + +@pytest.fixture +def db(tmp_path) -> Iterator[Database]: + database = Database(tmp_path / "sigmaforge.duckdb") + database.init_schema() + yield database + database.close() + + +def test_config_store_roundtrip(db: Database) -> None: + store = ConfigStore(db) + store.set("services.qdrant.port", 6333) + store.set("logging.level", "DEBUG") + + assert store.get("services.qdrant.port") == 6333 + assert store.get("logging.level") == "DEBUG" + assert store.get("missing") is None + assert store.get("missing", "default") == "default" + assert store.get_all() == {"services.qdrant.port": 6333, "logging.level": "DEBUG"} + assert store.delete("services.qdrant.port") is True + assert store.get("services.qdrant.port") is None + + +def test_model_store_active_selection(db: Database) -> None: + store = ModelStore(db) + store.upsert("embedding", "small", "/models/small", active=False) + store.upsert("embedding", "large", "/models/large", active=True) + + assert store.get_active("embedding") == "large" + store.set_active("embedding", "small") + assert store.get_active("embedding") == "small" + assert store.list("embedding") == [ + store.get("embedding", "large"), + store.get("embedding", "small"), + ] + + +def test_task_store_status_updates(db: Database) -> None: + store = TaskStore(db) + store.create(Task(id="t1", type="ingest", status="running", payload={"source": "local"})) + + task = store.get("t1") + assert task is not None + assert task.status == "running" + assert task.payload == {"source": "local"} + + assert store.update_status("t1", "completed", progress=1.0, message="done") is True + updated = store.get("t1") + assert updated is not None + assert updated.status == "completed" + assert updated.progress == 1.0 + assert updated.message == "done" + assert store.update_status("missing", "done") is False + assert len(store.list()) == 1 + + +def test_doc_registry_marks(db: Database) -> None: + registry = DocRegistry(db) + registry.mark("local", "rules/win.yaml", "indexed", indexed_at="2026-01-01T00:00:00") + + record = registry.get("local", "rules/win.yaml") + assert record is not None + assert record.status == "indexed" + assert record.indexed_at == "2026-01-01T00:00:00" + assert len(registry.list_records()) == 1 + assert len(registry.list_records(status="indexed")) == 1 + assert registry.delete("local", "rules/win.yaml") is True + assert registry.get("local", "rules/win.yaml") is None diff --git a/tests/rewrite/test_embed.py b/tests/rewrite/test_embed.py new file mode 100644 index 0000000..447aebb --- /dev/null +++ b/tests/rewrite/test_embed.py @@ -0,0 +1,49 @@ +from __future__ import annotations + +from sigmaforge.embed import Bm25SparseEncoder, DenseEncoder + + +class FakeModel: + def __init__(self, dim: int) -> None: + self.dim = dim + + def encode( + self, + texts, + batch_size: int | None = None, + normalize_embeddings: bool = False, + show_progress_bar: bool = False, + ): + return [[0.1] * self.dim for _ in texts] + + def get_sentence_embedding_dimension(self) -> int: + return self.dim + + +def test_sparse_encoder_is_deterministic() -> None: + encoder = Bm25SparseEncoder() + first = encoder.encode_text("attack attack defense") + second = encoder.encode_text("attack attack defense") + + assert first == second + assert len(first.indices) == 2 + assert all(value > 0 for value in first.values) + + +def test_sparse_encoder_batches_texts() -> None: + encoder = Bm25SparseEncoder() + vectors = encoder(["attack", "defense", "the a"]) + + assert len(vectors) == 3 + assert vectors[0].indices + assert vectors[1].indices + assert vectors[2].indices == () + + +def test_dense_encoder_uses_injected_model() -> None: + encoder = DenseEncoder("fake-model", model=FakeModel(dim=4)) + + vectors = encoder.encode(["one", "two"]) + + assert vectors == [[0.1] * 4, [0.1] * 4] + assert encoder.dim == 4 diff --git a/tests/rewrite/test_ingestion.py b/tests/rewrite/test_ingestion.py new file mode 100644 index 0000000..6637b38 --- /dev/null +++ b/tests/rewrite/test_ingestion.py @@ -0,0 +1,331 @@ +from __future__ import annotations + +from collections.abc import Iterator +from pathlib import Path + +import pytest + +from sigmaforge.embed import Bm25SparseEncoder, DenseEncoder +from sigmaforge.ingestion import ( + Chunk, + IngestRequest, + IngestionPipeline, + SigmaRule, + chunk_rule, + flat_rule_text, + parse_sigma_rule, + rich_rule_chunks, + rule_metadata, + split_text, +) +from sigmaforge.qdrant import QdrantCollectionManager, QdrantConnection, QdrantStore + + +class FakeModel: + def __init__(self, dim: int) -> None: + self.dim = dim + + def encode( + self, + texts, + batch_size: int | None = None, + normalize_embeddings: bool = False, + show_progress_bar: bool = False, + ): + return [[0.1] * self.dim for _ in texts] + + def get_sentence_embedding_dimension(self) -> int: + return self.dim + + +@pytest.fixture +def store() -> Iterator[QdrantStore]: + conn = QdrantConnection(location=":memory:") + manager = QdrantCollectionManager( + conn, + vector_size=4, + collections=("sigma_rules", "sigma_docs"), + enable_hybrid=True, + ) + manager.ensure() + yield QdrantStore(manager) + conn.close() + + +@pytest.fixture +def pipeline(store: QdrantStore) -> IngestionPipeline: + return IngestionPipeline( + store, + DenseEncoder("fake-model", model=FakeModel(dim=4)), + Bm25SparseEncoder(), + ) + + +def _rule() -> SigmaRule: + return SigmaRule( + id="suspicious", + title="Suspicious Rule", + condition="selection of process_name", + logsource={"product": "windows", "category": "process_creation"}, + detection={"process_name": {"contains": "powershell"}}, + ) + + +def _write_rule(path: Path, rule_id: str = "test_rule") -> Path: + path.write_text( + f"id: {rule_id}\n" + "title: Test Rule\n" + "description: Detects suspicious activity.\n" + "condition: selection of process_name\n" + "logsource:\n" + " product: windows\n" + " category: process_creation\n" + "detection:\n" + " process_name:\n" + " contains: powershell\n", + encoding="utf-8", + ) + return path + + +def test_parse_sigma_rule(tmp_path: Path) -> None: + path = _write_rule(tmp_path / "rule.yaml") + + rule = parse_sigma_rule(path) + + assert rule is not None + assert rule.id == "test_rule" + assert rule.title == "Test Rule" + assert rule.file_path == str(path) + assert rule.logsource["product"] == "windows" + + +def test_parse_sigma_rule_rejects_non_rule(tmp_path: Path) -> None: + path = tmp_path / "not-rule.yaml" + path.write_text("title: not a rule\n", encoding="utf-8") + + assert parse_sigma_rule(path) is None + + +def test_parse_sigma_rule_rejects_non_yaml_extension(tmp_path: Path) -> None: + path = tmp_path / "rule.txt" + path.write_text("detection:\n process_name: 1\n", encoding="utf-8") + + assert parse_sigma_rule(path) is None + + +def test_flat_rule_text_contains_core_fields() -> None: + text = flat_rule_text(_rule()) + + assert "Title: Suspicious Rule" in text + assert "Rule ID: suspicious" in text + assert "Condition: selection of process_name" in text + assert "product=windows" in text + assert "process_name: contains powershell" in text + + +def test_rule_metadata_exposes_search_fields() -> None: + metadata = rule_metadata(_rule()) + + assert metadata["rule_id"] == "suspicious" + assert metadata["title"] == "Suspicious Rule" + assert metadata["product"] == "windows" + assert metadata["category"] == "process_creation" + + +def test_rich_rule_chunks_creates_unique_chunks() -> None: + chunks = rich_rule_chunks(_rule()) + chunk_types = [chunk.chunk_type for chunk in chunks] + + assert len(chunks) > 1 + assert "executive_summary" in chunk_types + assert any(chunk_type.startswith("detection_block:") for chunk_type in chunk_types) + assert len({chunk.point_id("sigma_rules") for chunk in chunks}) == len(chunks) + + +def test_chunk_rule_rejects_unknown_mode() -> None: + with pytest.raises(ValueError): + chunk_rule(_rule(), "bogus") + + +def test_split_text_handles_empty() -> None: + assert split_text(" ") == [] + + +def test_split_text_short_text_single_chunk() -> None: + assert split_text("hello world", chunk_size=100, chunk_overlap=10) == ["hello world"] + + +def test_split_text_long_text_respects_chunk_size() -> None: + text = " ".join(f"token{i}" for i in range(1000)) + + chunks = split_text(text, chunk_size=80, chunk_overlap=10) + + assert chunks + assert all(len(chunk) <= 80 for chunk in chunks) + assert chunks[0].startswith("token0") + assert "token999" in chunks[-1] + + +def test_split_text_invalid_parameters() -> None: + with pytest.raises(ValueError): + split_text("text", chunk_size=0) + with pytest.raises(ValueError): + split_text("text", chunk_size=10, chunk_overlap=10) + + +def test_index_chunks_uses_stable_ids_for_same_chunk_type(store: QdrantStore) -> None: + pipeline = IngestionPipeline( + store, + DenseEncoder("fake-model", model=FakeModel(dim=4)), + Bm25SparseEncoder(), + ) + chunks = [ + Chunk(text="one", source_file="source.txt", chunk_type="part", chunk_index=0), + Chunk(text="two", source_file="source.txt", chunk_type="part", chunk_index=1), + ] + + indexed = pipeline.index_chunks("sigma_docs", chunks) + + assert indexed == 2 + assert store.manager.count("sigma_docs") == 2 + + +def test_ingest_sigma_directory_indexes_flat_rule( + tmp_path: Path, + store: QdrantStore, + pipeline: IngestionPipeline, +) -> None: + _write_rule(tmp_path / "rule.yaml") + + summary = pipeline.ingest_sigma_directory( + IngestRequest(directory=str(tmp_path)), + collection="sigma_rules", + ) + + assert summary.total == 1 + assert summary.indexed == 1 + assert summary.failed == 0 + assert summary.chunks == 1 + assert store.manager.count("sigma_rules") == 1 + + query_vector = pipeline.dense_encoder.encode(["powershell"])[0] + hits = store.query_dense("sigma_rules", query_vector, limit=1) + assert hits[0].payload["rule_id"] == "test_rule" + assert hits[0].payload["chunk_type"] == "rule" + assert hits[0].payload["source_file"].endswith("rule.yaml") + + +def test_ingest_sigma_directory_is_idempotent( + tmp_path: Path, + store: QdrantStore, + pipeline: IngestionPipeline, +) -> None: + _write_rule(tmp_path / "rule.yaml") + + first = pipeline.ingest_sigma_directory( + IngestRequest(directory=str(tmp_path)), + collection="sigma_rules", + ) + second = pipeline.ingest_sigma_directory( + IngestRequest(directory=str(tmp_path)), + collection="sigma_rules", + ) + + assert first.chunks == second.chunks == 1 + assert store.manager.count("sigma_rules") == 1 + + +def test_ingest_sigma_directory_rich_mode( + tmp_path: Path, + store: QdrantStore, + pipeline: IngestionPipeline, +) -> None: + _write_rule(tmp_path / "rule.yaml") + + summary = pipeline.ingest_sigma_directory( + IngestRequest(directory=str(tmp_path), mode="rich"), + collection="sigma_rules", + ) + + assert summary.chunks > 1 + assert store.manager.count("sigma_rules") == summary.chunks + + query_vector = pipeline.dense_encoder.encode(["powershell"])[0] + hits = store.query_dense("sigma_rules", query_vector, limit=summary.chunks) + assert len(hits) == summary.chunks + chunk_types = {hit.payload["chunk_type"] for hit in hits} + assert "executive_summary" in chunk_types + assert any(chunk_type.startswith("detection_block:") for chunk_type in chunk_types) + + +def test_ingest_sigma_directory_reports_invalid_rule( + tmp_path: Path, + store: QdrantStore, + pipeline: IngestionPipeline, +) -> None: + (tmp_path / "invalid.yaml").write_text("title: not a rule\n", encoding="utf-8") + + summary = pipeline.ingest_sigma_directory( + IngestRequest(directory=str(tmp_path)), + collection="sigma_rules", + ) + + assert summary.total == 1 + assert summary.failed == 1 + assert summary.indexed == 0 + assert store.manager.count("sigma_rules") == 0 + assert summary.results[0].error is not None + + +def test_ingest_sigma_directory_respects_selected_dirs( + tmp_path: Path, + store: QdrantStore, + pipeline: IngestionPipeline, +) -> None: + selected = tmp_path / "selected" + other = tmp_path / "other" + selected.mkdir() + other.mkdir() + _write_rule(selected / "a.yaml", rule_id="a") + _write_rule(other / "b.yaml", rule_id="b") + + summary = pipeline.ingest_sigma_directory( + IngestRequest(directory=str(tmp_path), selected_dirs=["selected"]), + collection="sigma_rules", + ) + + assert summary.total == 1 + assert summary.indexed == 1 + query_vector = pipeline.dense_encoder.encode(["powershell"])[0] + hits = store.query_dense("sigma_rules", query_vector, limit=1) + assert hits[0].payload["rule_id"] == "a" + + +def test_ingest_text_directory_chunks_and_indexes( + tmp_path: Path, + store: QdrantStore, +) -> None: + pipeline = IngestionPipeline( + store, + DenseEncoder("fake-model", model=FakeModel(dim=4)), + Bm25SparseEncoder(), + chunk_size=60, + chunk_overlap=10, + ) + (tmp_path / "doc.md").write_text("alpha beta gamma delta " * 20, encoding="utf-8") + + summary = pipeline.ingest_text_directory( + IngestRequest(directory=str(tmp_path)), + collection="sigma_docs", + ) + + assert summary.total == 1 + assert summary.indexed == 1 + assert summary.chunks > 1 + assert store.manager.count("sigma_docs") == summary.chunks + + query_vector = pipeline.dense_encoder.encode(["alpha"])[0] + hits = store.query_dense("sigma_docs", query_vector, limit=summary.chunks) + assert hits[0].payload["chunk_type"] == "doc" + assert hits[0].payload["source_file"].endswith("doc.md") diff --git a/tests/rewrite/test_llm.py b/tests/rewrite/test_llm.py new file mode 100644 index 0000000..c0db223 --- /dev/null +++ b/tests/rewrite/test_llm.py @@ -0,0 +1,79 @@ +from __future__ import annotations + +import asyncio + +import httpx +import pytest + +from sigmaforge.errors import LlamaError +from sigmaforge.llm import LlamaClient + + +def test_chat_returns_content() -> None: + def handler(request: httpx.Request) -> httpx.Response: + assert request.url.path == "/v1/chat/completions" + return httpx.Response(200, json={"choices": [{"message": {"content": "hello"}}]}) + + client = LlamaClient( + "http://llm.test", + transport=httpx.MockTransport(handler), + ) + result = asyncio.run( + client.chat( + [{"role": "user", "content": "hi"}], + ) + ) + + assert result == "hello" + + +def test_complete_returns_text() -> None: + def handler(request: httpx.Request) -> httpx.Response: + assert request.url.path == "/v1/completions" + return httpx.Response(200, json={"choices": [{"text": "completed"}]}) + + client = LlamaClient( + "http://llm.test", + transport=httpx.MockTransport(handler), + ) + + result = asyncio.run(client.complete("prompt")) + + assert result == "completed" + + +def test_stream_chat_yields_chunks() -> None: + body = ( + b'data: {"choices": [{"delta": {"content": "hel"}}]}\n' + b'data: {"choices": [{"delta": {"content": "lo"}}]}\n' + b"data: [DONE]\n" + ) + + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response(200, content=body, headers={"content-type": "text/event-stream"}) + + client = LlamaClient( + "http://llm.test", + transport=httpx.MockTransport(handler), + ) + + async def collect() -> list[str]: + chunks: list[str] = [] + async for chunk in client.stream_chat([{"role": "user", "content": "hi"}]): + chunks.append(chunk) + return chunks + + assert asyncio.run(collect()) == ["hel", "lo"] + + +def test_http_error_raises_llama_error() -> None: + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response(500, json={"detail": "boom"}) + + client = LlamaClient( + "http://llm.test", + transport=httpx.MockTransport(handler), + ) + + with pytest.raises(LlamaError): + asyncio.run(client.chat([{"role": "user", "content": "hi"}])) diff --git a/tests/rewrite/test_qdrant.py b/tests/rewrite/test_qdrant.py new file mode 100644 index 0000000..5451f97 --- /dev/null +++ b/tests/rewrite/test_qdrant.py @@ -0,0 +1,86 @@ +from __future__ import annotations + +from typing import Iterator + +import pytest + +from sigmaforge.qdrant import QdrantCollectionManager, QdrantConnection, QdrantStore + + +@pytest.fixture +def connection() -> Iterator[QdrantConnection]: + conn = QdrantConnection(location=":memory:") + yield conn + conn.close() + + +@pytest.fixture +def store(connection: QdrantConnection) -> QdrantStore: + manager = QdrantCollectionManager( + connection, + vector_size=4, + collections=("sigma_docs",), + enable_hybrid=True, + ) + manager.ensure() + return QdrantStore(manager) + + +def test_manager_creates_collection(store: QdrantStore) -> None: + assert store.manager.exists("sigma_docs") + assert store.manager.count("sigma_docs") == 0 + info = store.manager.collection_info("sigma_docs") + assert info["name"] == "sigma_docs" + + +def test_store_upsert_and_query(store: QdrantStore) -> None: + texts = ["windows powershell command execution"] + dense_vectors = [[0.1, 0.2, 0.3, 0.4]] + payloads = [{"source_file": "win.yaml", "chunk_type": "rule", "source": "local"}] + sparse_vectors = [((0, 1), (0.1, 0.2))] + + store.upsert_texts( + collection="sigma_docs", + texts=texts, + dense_vectors=dense_vectors, + payloads=payloads, + sparse_vectors=sparse_vectors, + ) + + assert store.manager.count("sigma_docs") == 1 + dense_hits = store.query_dense("sigma_docs", [0.1, 0.2, 0.3, 0.4], limit=1) + assert len(dense_hits) == 1 + assert dense_hits[0].payload["text"] == texts[0] + assert dense_hits[0].payload["source_file"] == "win.yaml" + + sparse_hits = store.query_sparse("sigma_docs", [0, 1], [0.1, 0.2], limit=1) + assert len(sparse_hits) == 1 + assert sparse_hits[0].id == dense_hits[0].id + + +def test_store_delete_by_source(store: QdrantStore) -> None: + store.upsert_texts( + collection="sigma_docs", + texts=["one", "two"], + dense_vectors=[[0.1, 0.2, 0.3, 0.4], [0.4, 0.3, 0.2, 0.1]], + payloads=[ + {"source_file": "win.yaml", "chunk_type": "rule"}, + {"source_file": "win.yaml", "chunk_type": "spec"}, + ], + ) + + store.delete_by_source("sigma_docs", "win.yaml") + + assert store.manager.count("sigma_docs") == 0 + + +def test_point_id_is_deterministic() -> None: + payload = {"source_file": "win.yaml", "chunk_type": "rule"} + first = QdrantStore.make_point_id("sigma_docs", payload) + second = QdrantStore.make_point_id("sigma_docs", payload) + other = QdrantStore.make_point_id( + "sigma_docs", {"source_file": "win.yaml", "chunk_type": "spec"} + ) + + assert first == second + assert first != other diff --git a/tests/rewrite/test_rag.py b/tests/rewrite/test_rag.py new file mode 100644 index 0000000..cbbe40a --- /dev/null +++ b/tests/rewrite/test_rag.py @@ -0,0 +1,259 @@ +from __future__ import annotations + +from collections.abc import AsyncIterator, Mapping, Sequence +from typing import Any + +from sigmaforge.errors import LlamaError +from sigmaforge.llm import LlamaClient +from sigmaforge.rag import RAGAnswer, RAGPipeline, render_search_prompt +from sigmaforge.search import SearchEngine, SearchResult + + +def _result(index: int, text: str, source: str) -> SearchResult: + return SearchResult( + id=f"id-{index}", + collection="sigma_rules", + text=text, + score=1.0 / index, + payload={"source_file": source}, + ) + + +class FakeSearchEngine: + def __init__(self, results: Sequence[SearchResult]) -> None: + self._results = list(results) + self.calls: list[tuple[str, int | None]] = [] + + def search( + self, + query: str, + *, + top_k: int | None = None, + collections: Sequence[str] | None = None, + filters: Mapping[str, str] | None = None, + extra_filter: Any | None = None, + ) -> list[SearchResult]: + self.calls.append((query, top_k)) + if top_k is not None: + return self._results[:top_k] + return list(self._results) + + +class FakeLlamaClient: + def __init__( + self, + response: str = "Default answer.", + chunks: Sequence[str] | None = None, + *, + fail_chat: bool = False, + fail_stream: bool = False, + ) -> None: + self.model = "fake-model" + self._response = response + self._chunks = list(chunks) if chunks is not None else [response] + self._fail_chat = fail_chat + self._fail_stream = fail_stream + self.messages: list[dict[str, Any]] = [] + self.stream_messages: list[dict[str, Any]] = [] + + async def chat( + self, + messages: Sequence[Mapping[str, Any]], + *, + temperature: float = 0.3, + max_tokens: int = 512, + stop: str | Sequence[str] | None = None, + ) -> str: + self.messages = [dict(message) for message in messages] + if self._fail_chat: + raise LlamaError("chat failed") + return self._response + + async def stream_chat( + self, + messages: Sequence[Mapping[str, Any]], + *, + temperature: float = 0.3, + max_tokens: int = 512, + stop: str | Sequence[str] | None = None, + ) -> AsyncIterator[str]: + self.stream_messages = [dict(message) for message in messages] + if self._fail_stream: + raise LlamaError("stream failed") + for chunk in self._chunks: + yield chunk + + +class PartialFailLlamaClient: + model = "fake-model" + + async def chat( + self, + messages: Sequence[Mapping[str, Any]], + *, + temperature: float = 0.3, + max_tokens: int = 512, + stop: str | Sequence[str] | None = None, + ) -> str: + return "unused" + + async def stream_chat( + self, + messages: Sequence[Mapping[str, Any]], + *, + temperature: float = 0.3, + max_tokens: int = 512, + stop: str | Sequence[str] | None = None, + ) -> AsyncIterator[str]: + yield "partial" + raise LlamaError("stream failed after partial output") + + +def test_render_search_prompt_includes_context_and_question() -> None: + prompt = render_search_prompt("search context", "question text") + assert "search context" in prompt + assert "question text" in prompt + + +def test_real_search_and_llm_types_are_accepted_by_pipeline() -> None: + def create_pipeline(search_engine: SearchEngine, llm_client: LlamaClient) -> RAGPipeline: + return RAGPipeline(search_engine, llm_client) + + assert create_pipeline is not None + + +def test_rag_answer_to_dict_includes_sources() -> None: + result = _result(1, "rule alpha", "alpha.yml") + answer = RAGAnswer( + query="q", + answer="a", + sources=(result,), + model="m", + fallback=False, + ) + + data = answer.to_dict() + + assert data["query"] == "q" + assert data["answer"] == "a" + assert data["model"] == "m" + assert data["fallback"] is False + assert data["sources"][0]["id"] == "id-1" + + +async def test_answer_returns_llm_response_and_sources() -> None: + results = [ + _result(1, "rule alpha detects powershell", "alpha.yml"), + _result(2, "rule beta detects wmi", "beta.yml"), + ] + search = FakeSearchEngine(results) + llm = FakeLlamaClient("Rule alpha detects PowerShell.") + pipeline = RAGPipeline(search, llm) + + answer = await pipeline.answer("which rule detects powershell?") + + assert isinstance(answer, RAGAnswer) + assert answer.answer == "Rule alpha detects PowerShell." + assert answer.fallback is False + assert answer.model == "fake-model" + assert answer.sources == tuple(results) + assert search.calls == [("which rule detects powershell?", 5)] + + system = str(llm.messages[0]["content"]) + user = str(llm.messages[1]["content"]) + assert "which rule detects powershell?" in system + assert "rule alpha detects powershell" in system + assert user == "which rule detects powershell?" + + +async def test_answer_passes_top_k_override_to_search() -> None: + results = [ + _result(1, "rule alpha", "alpha.yml"), + _result(2, "rule beta", "beta.yml"), + _result(3, "rule gamma", "gamma.yml"), + ] + search = FakeSearchEngine(results) + llm = FakeLlamaClient() + pipeline = RAGPipeline(search, llm) + + answer = await pipeline.answer("query", top_k=2) + + assert search.calls == [("query", 2)] + assert len(answer.sources) == 2 + + +async def test_answer_without_results_returns_fallback_without_llm_call() -> None: + search = FakeSearchEngine([]) + llm = FakeLlamaClient() + pipeline = RAGPipeline(search, llm) + + answer = await pipeline.answer("missing") + + assert answer.fallback is True + assert answer.sources == () + assert "No matching" in answer.answer + assert llm.messages == [] + + +async def test_answer_llm_failure_returns_fallback_with_sources() -> None: + results = [ + _result(1, "rule alpha detects powershell", "alpha.yml"), + _result(2, "rule beta detects wmi", "beta.yml"), + ] + search = FakeSearchEngine(results) + llm = FakeLlamaClient(fail_chat=True) + pipeline = RAGPipeline(search, llm) + + answer = await pipeline.answer("which rule detects powershell?") + + assert answer.fallback is True + assert answer.sources == tuple(results) + assert "LLM unavailable" in answer.answer + assert "rule alpha detects powershell" in answer.answer + + +async def test_answer_stream_yields_llm_chunks() -> None: + results = [_result(1, "rule alpha detects powershell", "alpha.yml")] + search = FakeSearchEngine(results) + llm = FakeLlamaClient(chunks=["Pow", "erShell"]) + pipeline = RAGPipeline(search, llm) + + chunks = [chunk async for chunk in pipeline.answer_stream("powershell")] + + assert chunks == ["Pow", "erShell"] + system = str(llm.stream_messages[0]["content"]) + assert "powershell" in system + assert "rule alpha detects powershell" in system + + +async def test_answer_stream_without_results_yields_fallback() -> None: + search = FakeSearchEngine([]) + llm = FakeLlamaClient() + pipeline = RAGPipeline(search, llm) + + chunks = [chunk async for chunk in pipeline.answer_stream("missing")] + + assert chunks == ["No matching Sigma rules or documents were found."] + + +async def test_answer_stream_failure_before_output_yields_fallback() -> None: + results = [_result(1, "rule alpha detects powershell", "alpha.yml")] + search = FakeSearchEngine(results) + llm = FakeLlamaClient(fail_stream=True) + pipeline = RAGPipeline(search, llm) + + chunks = [chunk async for chunk in pipeline.answer_stream("powershell")] + + assert len(chunks) == 1 + assert "LLM unavailable" in chunks[0] + + +async def test_answer_stream_failure_after_partial_output_does_not_add_fallback() -> None: + results = [_result(1, "rule alpha detects powershell", "alpha.yml")] + search = FakeSearchEngine(results) + llm = PartialFailLlamaClient() + pipeline = RAGPipeline(search, llm) + + chunks = [chunk async for chunk in pipeline.answer_stream("powershell")] + + assert chunks == ["partial"] diff --git a/tests/rewrite/test_search.py b/tests/rewrite/test_search.py new file mode 100644 index 0000000..d407b49 --- /dev/null +++ b/tests/rewrite/test_search.py @@ -0,0 +1,501 @@ +from __future__ import annotations + +from collections.abc import Iterator, Sequence +from typing import Any + +import pytest +from qdrant_client import models + +from sigmaforge.embed import Bm25SparseEncoder +from sigmaforge.errors import SearchError +from sigmaforge.qdrant import QdrantCollectionManager, QdrantConnection, QdrantStore +from sigmaforge.search import ( + SearchEngine, + SearchResult, + build_qdrant_filter, + format_context, + format_result, + get_citation, + parse_query_filters, + reciprocal_rank_fusion, +) + + +class FixedDenseEncoder: + def __init__(self, vector: list[float]) -> None: + self.vector = vector + self.calls: list[str] = [] + + def encode(self, texts: Sequence[str], *, is_query: bool = False) -> list[list[float]]: + self.calls.extend(texts) + return [list(self.vector) for _ in texts] + + +@pytest.fixture +def connection() -> Iterator[QdrantConnection]: + conn = QdrantConnection(location=":memory:") + yield conn + conn.close() + + +@pytest.fixture +def store(connection: QdrantConnection) -> QdrantStore: + manager = QdrantCollectionManager( + connection, + vector_size=4, + collections=("sigma_rules", "sigma_docs"), + enable_hybrid=True, + ) + manager.ensure() + return QdrantStore(manager) + + +def _result( + result_id: str, + collection: str = "sigma_rules", + text: str = "text", + score: float = 1.0, + payload: dict[str, Any] | None = None, +) -> SearchResult: + return SearchResult( + id=result_id, + collection=collection, + text=text, + score=score, + payload=payload if payload is not None else {"source_file": f"{result_id}.yaml"}, + ) + + +def _upsert_point( + store: QdrantStore, + collection: str, + text: str, + payload: dict[str, Any], + dense: list[float] | None = None, +) -> None: + sparse_encoder = Bm25SparseEncoder() + vector = sparse_encoder.encode_text(text) + store.upsert_texts( + collection=collection, + texts=[text], + dense_vectors=[dense or [1.0, 0.0, 0.0, 0.0]], + payloads=[payload], + sparse_vectors=[(vector.indices, vector.values)], + ) + + +def test_parse_query_filters_extracts_known_filters() -> None: + filters, cleaned = parse_query_filters("powershell rule_id:alpha level:high") + + assert filters == {"rule_id": "alpha", "level": "high"} + assert cleaned == "powershell" + + +def test_parse_query_filters_keeps_unknown_colons() -> None: + filters, cleaned = parse_query_filters("process_name:powershell rule_id:alpha") + + assert filters == {"rule_id": "alpha"} + assert cleaned == "process_name:powershell" + + +def test_parse_query_filters_returns_empty_cleaned_for_only_filters() -> None: + filters, cleaned = parse_query_filters("rule_id:alpha") + + assert filters == {"rule_id": "alpha"} + assert cleaned == "" + + +def test_build_qdrant_filter_returns_none_without_filters() -> None: + assert build_qdrant_filter() is None + assert build_qdrant_filter({}) is None + + +def test_build_qdrant_filter_uses_scalar_match_value() -> None: + qfilter = build_qdrant_filter({"rule_id": "alpha"}) + + assert qfilter is not None + data = qfilter.model_dump() + condition = data["must"][0] + assert condition["key"] == "rule_id" + assert condition["match"]["value"] == "alpha" + + +def test_build_qdrant_filter_uses_match_any_for_list_filters() -> None: + qfilter = build_qdrant_filter({"tags": "alpha beta"}) + + assert qfilter is not None + data = qfilter.model_dump() + condition = data["must"][0] + assert condition["key"] == "tags" + assert condition["match"]["any"] == ["alpha", "beta"] + + +def test_build_qdrant_filter_combines_extra_filter() -> None: + extra = models.Filter( + must=[models.FieldCondition(key="status", match=models.MatchValue(value="active"))] + ) + + qfilter = build_qdrant_filter({"rule_id": "alpha"}, extra) + + assert qfilter is not None + data = qfilter.model_dump() + assert isinstance(data["must"], list) + assert len(data["must"]) == 2 + + +def test_reciprocal_rank_fusion_merges_rankings() -> None: + first = _result("1") + second = _result("2") + duplicate = _result("1", score=0.5) + + fused = reciprocal_rank_fusion({"dense": [first, second], "sparse": [duplicate]}) + + assert [result.id for result in fused] == ["1", "2"] + assert fused[0].score == pytest.approx(2 / 61) + assert first.score == 1.0 + + +def test_reciprocal_rank_fusion_ignores_non_positive_weights() -> None: + fused = reciprocal_rank_fusion( + {"dense": [_result("1")], "sparse": [_result("2")]}, + weights={"dense": 1.0, "sparse": 0.0}, + ) + + assert [result.id for result in fused] == ["1"] + + +def test_reciprocal_rank_fusion_limits_results() -> None: + fused = reciprocal_rank_fusion({"dense": [_result(str(i)) for i in range(5)]}, limit=2) + + assert len(fused) == 2 + + +def test_reciprocal_rank_fusion_rejects_non_positive_k() -> None: + with pytest.raises(SearchError): + reciprocal_rank_fusion({"dense": []}, k=0) + + +def test_format_context_orders_and_labels_results() -> None: + results = [ + _result("1", text="alpha", payload={"source_file": "a.yaml", "title": "Alpha"}), + _result("2", collection="sigma_docs", text="beta", payload={"source_file": "b.md"}), + ] + + context = format_context(results, max_results=1) + + assert context.startswith("[1] sigma_rules / a.yaml / Alpha") + assert "alpha" in context + assert "[2]" not in context + + +def test_format_context_respects_budget() -> None: + result = _result("1", text="a" * 100) + + context = format_context([result], max_chars=50) + + assert len(context) <= 50 + assert context.endswith("...") + + +def test_get_citation_uses_source_and_line() -> None: + result = _result("1", payload={"source_file": "a.yaml", "line_start": 10}) + + assert get_citation(result) == "a.yaml:10" + + result_without_line = _result("1", payload={"source_file": "a.yaml"}) + assert get_citation(result_without_line) == "a.yaml" + + +def test_format_result_exposes_collection_specific_fields() -> None: + rule = _result( + "1", + text="rule text", + payload={"source_file": "a.yaml", "rule_id": "r", "title": "T", "level": "high"}, + ) + formatted = format_result(rule) + + assert formatted["rule_id"] == "r" + assert formatted["title"] == "T" + + doc = _result( + "2", + collection="sigma_docs", + text="doc text", + payload={"source_file": "b.md", "doc_type": "doc", "original_url": "https://example.com"}, + ) + formatted_doc = format_result(doc) + + assert formatted_doc["doc_type"] == "doc" + assert formatted_doc["original_url"] == "https://example.com" + + +def test_search_empty_query_returns_empty_and_does_not_encode(store: QdrantStore) -> None: + encoder = FixedDenseEncoder([1.0, 0.0, 0.0, 0.0]) + engine = SearchEngine(store, encoder) + + assert engine.search("") == [] + assert engine.search("???") == [] + assert encoder.calls == [] + + +def test_search_normalizes_query(store: QdrantStore) -> None: + encoder = FixedDenseEncoder([1.0, 0.0, 0.0, 0.0]) + _upsert_point( + store, + "sigma_rules", + "powershell process creation", + {"source_file": "ps.yaml", "chunk_type": "rule", "rule_id": "ps"}, + ) + engine = SearchEngine(store, encoder, alpha=1.0) + + results = engine.search("`powershell`?") + + assert len(results) == 1 + assert encoder.calls == ["powershell"] + + +def test_search_falls_back_to_normalized_query_when_only_filters( + store: QdrantStore, +) -> None: + encoder = FixedDenseEncoder([1.0, 0.0, 0.0, 0.0]) + _upsert_point( + store, + "sigma_rules", + "powershell process creation", + {"source_file": "ps.yaml", "chunk_type": "rule", "rule_id": "alpha"}, + ) + engine = SearchEngine(store, encoder, alpha=1.0) + + results = engine.search("rule_id:alpha") + + assert len(results) == 1 + assert results[0].metadata["rule_id"] == "alpha" + assert encoder.calls == ["rule_id:alpha"] + + +def test_search_returns_filtered_results(store: QdrantStore) -> None: + encoder = FixedDenseEncoder([1.0, 0.0, 0.0, 0.0]) + _upsert_point( + store, + "sigma_rules", + "powershell process creation", + { + "source_file": "alpha.yaml", + "chunk_type": "rule", + "rule_id": "alpha", + "status": "active", + }, + ) + _upsert_point( + store, + "sigma_rules", + "wmi process creation", + { + "source_file": "beta.yaml", + "chunk_type": "rule", + "rule_id": "beta", + "status": "disabled", + }, + ) + engine = SearchEngine(store, encoder, alpha=1.0) + + results = engine.search("powershell rule_id:alpha", top_k=5) + + assert len(results) == 1 + assert results[0].metadata["rule_id"] == "alpha" + assert encoder.calls == ["powershell"] + + +def test_search_extra_filter_applies_to_results(store: QdrantStore) -> None: + encoder = FixedDenseEncoder([1.0, 0.0, 0.0, 0.0]) + _upsert_point( + store, + "sigma_rules", + "powershell process creation", + { + "source_file": "alpha.yaml", + "chunk_type": "rule", + "rule_id": "alpha", + "status": "active", + }, + ) + _upsert_point( + store, + "sigma_rules", + "wmi process creation", + { + "source_file": "beta.yaml", + "chunk_type": "rule", + "rule_id": "beta", + "status": "disabled", + }, + ) + engine = SearchEngine(store, encoder, alpha=1.0) + extra_filter = models.Filter( + must=[models.FieldCondition(key="status", match=models.MatchValue(value="active"))] + ) + + results = engine.search("powershell", extra_filter=extra_filter) + + assert len(results) == 1 + assert results[0].metadata["rule_id"] == "alpha" + + +def test_search_list_filter_matches_any_tag(store: QdrantStore) -> None: + encoder = FixedDenseEncoder([1.0, 0.0, 0.0, 0.0]) + _upsert_point( + store, + "sigma_rules", + "powershell process creation", + { + "source_file": "alpha.yaml", + "chunk_type": "rule", + "rule_id": "alpha", + "tags": ["powershell", "windows"], + }, + ) + _upsert_point( + store, + "sigma_rules", + "wmi process creation", + { + "source_file": "beta.yaml", + "chunk_type": "rule", + "rule_id": "beta", + "tags": ["wmi"], + }, + ) + engine = SearchEngine(store, encoder, alpha=1.0) + + results = engine.search("powershell tags:powershell") + + assert len(results) == 1 + assert results[0].metadata["rule_id"] == "alpha" + + +def test_search_references_filter_includes_sigma_docs(store: QdrantStore) -> None: + encoder = FixedDenseEncoder([1.0, 0.0, 0.0, 0.0]) + _upsert_point( + store, + "sigma_rules", + "powershell process creation", + {"source_file": "alpha.yaml", "chunk_type": "rule", "rule_id": "alpha"}, + ) + _upsert_point( + store, + "sigma_docs", + "powershell documentation", + { + "source_file": "doc.md", + "chunk_type": "doc", + "references": "docref", + }, + ) + engine = SearchEngine(store, encoder, alpha=1.0) + + results = engine.search("powershell references:docref", collections=("sigma_rules",)) + + assert len(results) == 1 + assert results[0].collection == "sigma_docs" + + +def test_search_collection_queries_single_collection(store: QdrantStore) -> None: + encoder = FixedDenseEncoder([1.0, 0.0, 0.0, 0.0]) + _upsert_point( + store, + "sigma_rules", + "powershell process creation", + {"source_file": "alpha.yaml", "chunk_type": "rule", "rule_id": "alpha"}, + ) + _upsert_point( + store, + "sigma_docs", + "powershell documentation", + {"source_file": "doc.md", "chunk_type": "doc", "references": "docref"}, + ) + engine = SearchEngine(store, encoder, alpha=1.0) + + results = engine.search_collection("sigma_rules", "powershell") + + assert len(results) == 1 + assert results[0].collection == "sigma_rules" + + +def test_search_sparse_only_returns_matching_document(store: QdrantStore) -> None: + encoder = FixedDenseEncoder([1.0, 0.0, 0.0, 0.0]) + _upsert_point( + store, + "sigma_rules", + "powershell process creation", + { + "source_file": "alpha.yaml", + "chunk_type": "rule", + "rule_id": "powershell", + }, + dense=[0.0, 1.0, 0.0, 0.0], + ) + _upsert_point( + store, + "sigma_rules", + "wmi event", + {"source_file": "beta.yaml", "chunk_type": "rule", "rule_id": "wmi"}, + dense=[0.0, 1.0, 0.0, 0.0], + ) + engine = SearchEngine(store, encoder, Bm25SparseEncoder(), alpha=0.0) + + results = engine.search("powershell", collections=("sigma_rules",)) + + assert len(results) == 1 + assert results[0].metadata["rule_id"] == "powershell" + + +def test_search_sparse_only_skips_stopword_only_query(store: QdrantStore) -> None: + encoder = FixedDenseEncoder([1.0, 0.0, 0.0, 0.0]) + _upsert_point( + store, + "sigma_rules", + "powershell process creation", + {"source_file": "alpha.yaml", "chunk_type": "rule", "rule_id": "alpha"}, + ) + engine = SearchEngine(store, encoder, Bm25SparseEncoder(), alpha=0.0) + + results = engine.search("the", collections=("sigma_rules",)) + + assert results == [] + + +def test_search_similarity_threshold_filters_low_scores(store: QdrantStore) -> None: + encoder = FixedDenseEncoder([0.0, 1.0, 0.0, 0.0]) + _upsert_point( + store, + "sigma_rules", + "powershell process creation", + {"source_file": "alpha.yaml", "chunk_type": "rule", "rule_id": "alpha"}, + dense=[1.0, 0.0, 0.0, 0.0], + ) + engine = SearchEngine(store, encoder, alpha=1.0, similarity_threshold=0.5) + + results = engine.search("powershell", collections=("sigma_rules",)) + + assert results == [] + + +def test_search_top_k_limits_results(store: QdrantStore) -> None: + encoder = FixedDenseEncoder([1.0, 0.0, 0.0, 0.0]) + _upsert_point( + store, + "sigma_rules", + "powershell process creation", + {"source_file": "alpha.yaml", "chunk_type": "rule", "rule_id": "alpha"}, + ) + _upsert_point( + store, + "sigma_rules", + "wmi process creation", + {"source_file": "beta.yaml", "chunk_type": "rule", "rule_id": "beta"}, + ) + engine = SearchEngine(store, encoder, alpha=1.0) + + results = engine.search("powershell", top_k=1) + + assert len(results) == 1 diff --git a/uv.lock b/uv.lock index 0525084..1bda163 100644 --- a/uv.lock +++ b/uv.lock @@ -2007,6 +2007,25 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/75/a6/a0a304dc33b49145b21f4808d763822111e67d1c3a32b524a1baf947b6e1/platformdirs-4.9.6-py3-none-any.whl", hash = "sha256:e61adb1d5e5cb3441b4b7710bea7e4c12250ca49439228cc1021c00dcfac0917", size = 21348, upload-time = "2026-04-09T00:04:09.463Z" }, ] +[[package]] +name = "playwright" +version = "1.62.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "greenlet" }, + { name = "pyee" }, +] +wheels = [ + { url = "https://files.pythonhosted.org/packages/6c/5b/ca2abcf3aa69f9fb510215e3064f30b57fe57657c8d04ede45bb966d5606/playwright-1.62.0-py3-none-macosx_10_13_x86_64.whl", hash = "sha256:d8da938f3748841a8754f2e1f0216902c1c8f8ae3720de8b32ccf8e6913a7c4f", size = 43732091, upload-time = "2026-07-31T17:00:44.178Z" }, + { url = "https://files.pythonhosted.org/packages/af/1a/0bfbe9904350961f4dbb713f04342e40d548c5fc26c8157bd13617c81492/playwright-1.62.0-py3-none-macosx_11_0_arm64.whl", hash = "sha256:db755ab27db21a04186f1fe8169888e42356086e439b1059b923ef417f0b6034", size = 42510842, upload-time = "2026-07-31T17:00:48.596Z" }, + { url = "https://files.pythonhosted.org/packages/66/dc/c0486b407ad0699a250f6bbe3066fca95344009a99ca66e88ca175c69dc1/playwright-1.62.0-py3-none-macosx_11_0_universal2.whl", hash = "sha256:5108bd5b3e87169ddf269feee097da5893af7f8aea4634dfc840518d64c1f1da", size = 43732093, upload-time = "2026-07-31T17:00:52.218Z" }, + { url = "https://files.pythonhosted.org/packages/43/6b/b24aebc2b04bffcb342bccf96e287c78b363e1615bed5cea97500cc0393a/playwright-1.62.0-py3-none-manylinux1_x86_64.whl", hash = "sha256:ba33bae6a13b3d9d354c751cb618af357d20fe1d57767cbcce52079bbef17ad3", size = 47748926, upload-time = "2026-07-31T17:00:56.438Z" }, + { url = "https://files.pythonhosted.org/packages/36/43/b4b18bdc87e1949568fffdcde3ff9a0456266b2d0c6d4432cc34d89ea6eb/playwright-1.62.0-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:db2d76613a57ad844362ce42f7d0c2fa26b19a4f7a46d4f76b891c631e6e5aff", size = 47441423, upload-time = "2026-07-31T17:01:00.404Z" }, + { url = "https://files.pythonhosted.org/packages/81/22/af5d926fc2c32a339eec00a443644bc40ab9db1dd2dd9017873c59773c0c/playwright-1.62.0-py3-none-win32.whl", hash = "sha256:e5614fa89355d7081457680324bb219f79f69c423c5cb6fa250e30b0d8aebf1c", size = 38164450, upload-time = "2026-07-31T17:01:04.187Z" }, + { url = "https://files.pythonhosted.org/packages/2b/a9/4160c1033c07af98bf841ad079457dd78408a5ee0dd56cbfe50b8b6a1c22/playwright-1.62.0-py3-none-win_amd64.whl", hash = "sha256:92c0d98ed04eb35af557b709875edba415b1f548bdb22ddb5bb3e1e6c835c2f1", size = 38164458, upload-time = "2026-07-31T17:01:08.459Z" }, + { url = "https://files.pythonhosted.org/packages/6c/ec/06b55d619a7082a766aa04f2c6bb31435c87f02930087d8a0517119408fa/playwright-1.62.0-py3-none-win_arm64.whl", hash = "sha256:ea8d3055aa9d5a9f1832ac82517bd8b42c78fac7ebcbebb0107116735c8cb6a1", size = 34208868, upload-time = "2026-07-31T17:01:11.818Z" }, +] + [[package]] name = "pluggy" version = "1.6.0" @@ -2242,6 +2261,18 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/17/eb/9d89ad2d9b0ba8cd65393d434471621b98912abb10fbe1df08e480ba57b5/pydantic_core-2.46.3-graalpy312-graalpy250_312_native-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:fd35aa21299def8db7ef4fe5c4ff862941a9a158ca7b63d61e66fe67d30416b4", size = 2137657, upload-time = "2026-04-20T14:42:45.149Z" }, ] +[[package]] +name = "pyee" +version = "13.0.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/8b/04/e7c1fe4dc78a6fdbfd6c337b1c3732ff543b8a397683ab38378447baa331/pyee-13.0.1.tar.gz", hash = "sha256:0b931f7c14535667ed4c7e0d531716368715e860b988770fc7eb8578d1f67fc8", size = 31655, upload-time = "2026-02-14T21:12:28.044Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/a0/c4/b4d4827c93ef43c01f599ef31453ccc1c132b353284fc6c87d535c233129/pyee-13.0.1-py3-none-any.whl", hash = "sha256:af2f8fede4171ef667dfded53f96e2ed0d6e6bd7ee3bb46437f77e3b57689228", size = 15659, upload-time = "2026-02-14T21:12:26.263Z" }, +] + [[package]] name = "pygments" version = "2.20.0" @@ -2306,6 +2337,19 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/03/e2/08a497ef684b88559c9cc5f4ad53a37e7b99e727094a86d6ea32536d5d3c/pytest_asyncio-1.4.0-py3-none-any.whl", hash = "sha256:933ca923a23075a87fb7070c0ec272a6848489824d887c85c812670932835aa1", size = 16930, upload-time = "2026-05-26T09:56:02.576Z" }, ] +[[package]] +name = "pytest-base-url" +version = "2.1.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "pytest" }, + { name = "requests" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/ae/1a/b64ac368de6b993135cb70ca4e5d958a5c268094a3a2a4cac6f0021b6c4f/pytest_base_url-2.1.0.tar.gz", hash = "sha256:02748589a54f9e63fcbe62301d6b0496da0d10231b753e950c63e03aee745d45", size = 6702, upload-time = "2024-01-31T22:43:00.81Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/98/1c/b00940ab9eb8ede7897443b771987f2f4a76f06be02f1b3f01eb7567e24a/pytest_base_url-2.1.0-py3-none-any.whl", hash = "sha256:3ad15611778764d451927b2a53240c1a7a591b521ea44cebfe45849d2d2812e6", size = 5302, upload-time = "2024-01-31T22:42:58.897Z" }, +] + [[package]] name = "pytest-cov" version = "7.1.0" @@ -2320,6 +2364,21 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/9d/7a/d968e294073affff457b041c2be9868a40c1c71f4a35fcc1e45e5493067b/pytest_cov-7.1.0-py3-none-any.whl", hash = "sha256:a0461110b7865f9a271aa1b51e516c9a95de9d696734a2f71e3e78f46e1d4678", size = 22876, upload-time = "2026-03-21T20:11:14.438Z" }, ] +[[package]] +name = "pytest-playwright" +version = "0.9.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "playwright" }, + { name = "pytest" }, + { name = "pytest-base-url" }, + { name = "python-slugify" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/26/c4/31cab1a9edfa35546950bb30c033dbfe4d4f3a216f4bafb0537a1469a70c/pytest_playwright-0.9.0.tar.gz", hash = "sha256:bd44daa852b0fb8b0e55a2f88b727507dedda0e350bea5d09420647252f2454b", size = 17899, upload-time = "2026-08-10T10:20:16.407Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/ab/70/72ff5c3833f0faa9b5fb23f24c0a606f070ad33bb61b802ba38dfe125e7d/pytest_playwright-0.9.0-py3-none-any.whl", hash = "sha256:9d9dc74e335c647944cecfffa706cc9e6e4bf4c25e84592c128b584d494d4471", size = 17915, upload-time = "2026-08-10T10:20:18.253Z" }, +] + [[package]] name = "python-dateutil" version = "2.9.0.post0" @@ -2369,6 +2428,18 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/d9/4f/00be2196329ebbff56ce564aa94efb0fbc828d00de250b1980de1a34ab49/python_pptx-1.0.2-py3-none-any.whl", hash = "sha256:160838e0b8565a8b1f67947675886e9fea18aa5e795db7ae531606d68e785cba", size = 472788, upload-time = "2024-08-07T17:33:28.192Z" }, ] +[[package]] +name = "python-slugify" +version = "8.0.4" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "text-unidecode" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/87/c7/5e1547c44e31da50a460df93af11a535ace568ef89d7a811069ead340c4a/python-slugify-8.0.4.tar.gz", hash = "sha256:59202371d1d05b54a9e7720c5e038f928f45daaffe41dd10822f3907b937c856", size = 10921, upload-time = "2024-02-08T18:32:45.488Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/a4/62/02da182e544a51a5c3ccf4b03ab79df279f9c60c5e82d5e8bec7ca26ac11/python_slugify-8.0.4-py2.py3-none-any.whl", hash = "sha256:276540b79961052b66b7d116620b36518847f52d5fd9e3a70164fc8c50faa6b8", size = 10051, upload-time = "2024-02-08T18:32:43.911Z" }, +] + [[package]] name = "pytz" version = "2026.2" @@ -2797,10 +2868,12 @@ dependencies = [ [package.dev-dependencies] dev = [ { name = "mypy" }, + { name = "playwright" }, { name = "pre-commit" }, { name = "pytest" }, { name = "pytest-asyncio" }, { name = "pytest-cov" }, + { name = "pytest-playwright" }, { name = "ruff" }, { name = "types-pyyaml" }, ] @@ -2836,10 +2909,12 @@ requires-dist = [ [package.metadata.requires-dev] dev = [ { name = "mypy" }, + { name = "playwright", specifier = ">=1.48" }, { name = "pre-commit", specifier = ">=4.6.0" }, { name = "pytest", specifier = ">=9.0.3" }, { name = "pytest-asyncio", specifier = ">=1.4.0" }, { name = "pytest-cov", specifier = ">=7.1.0" }, + { name = "pytest-playwright", specifier = ">=0.6" }, { name = "ruff" }, { name = "types-pyyaml", specifier = ">=6.0.12.20260518" }, ] @@ -2974,6 +3049,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/d7/c1/eb8f9debc45d3b7918a32ab756658a0904732f75e555402972246b0b8e71/tenacity-9.1.4-py3-none-any.whl", hash = "sha256:6095a360c919085f28c6527de529e76a06ad89b23659fa881ae0649b867a9d55", size = 28926, upload-time = "2026-02-07T10:45:32.24Z" }, ] +[[package]] +name = "text-unidecode" +version = "1.3" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/ab/e2/e9a00f0ccb71718418230718b3d900e71a5d16e701a3dae079a21e9cd8f8/text-unidecode-1.3.tar.gz", hash = "sha256:bad6603bb14d279193107714b288be206cac565dfa49aa5b105294dd5c4aab93", size = 76885, upload-time = "2019-08-30T21:36:45.405Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/a6/a5/c0b6468d3824fe3fde30dbb5e1f687b291608f9473681bbf7dabbf5a87d7/text_unidecode-1.3-py2.py3-none-any.whl", hash = "sha256:1311f10e8b895935241623731c2ba64f4c455287888b18189350b67134a822e8", size = 78154, upload-time = "2019-08-30T21:37:03.543Z" }, +] + [[package]] name = "threadpoolctl" version = "3.6.0"