From a80b3639e0314b257186ae352fa1281fcb533a45 Mon Sep 17 00:00:00 2001 From: cbizeul Date: Fri, 29 May 2026 11:47:57 +0000 Subject: [PATCH 01/11] refactor: move logger and monitoring to core/, repoint imports Phase 12A. Relocate the cross-cutting observability utilities out of the old utils/ tree into the hexagonal core layer: - utils/logger.py -> core/utils/logging.py (structlog-convention name) - utils/monitoring.py -> core/observability/monitoring.py - utils/test_logger.py -> core/utils/test_logging.py logging.py now reads config via `from core.config import load_config`; the core.config package exposes the cached load_config/get_settings/Settings public API (previously only the config/ shim did). Repoint every import site to the new paths: - from utils.logger -> from core.utils.logging - from utils.monitoring -> from core.observability.monitoring - from utils.exceptions.base -> from core.utils.exceptions - from utils.exceptions.embeddings -> from core.utils.exceptions - from utils.external_resource_errors -> from core.utils.external_errors The exceptions and external_errors modules already live canonically in core.utils; this only switches the remaining consumers off the shims. Delete the duplicate utils/test_external_resource_errors.py (kept in core/utils/test_external_errors.py). Update the auth-router test stub to patch core.utils.logging instead of the moved utils.logger. No remaining `from utils.` imports outside the (now shim-only) utils/ dir. --- CLAUDE.md | 4 +- .../docs/documentation/milvus_migration.mdx | 2 +- openrag/api/dependencies/auth.py | 2 +- openrag/api/dependencies/llm.py | 2 +- openrag/api/error_handlers.py | 2 +- openrag/api/main.py | 2 +- openrag/api/mcp/server.py | 2 +- openrag/api/middleware/auth.py | 2 +- openrag/api/middleware/instrumentation.py | 2 +- openrag/api/routers/admin/cluster.py | 2 +- openrag/api/routers/admin/indexing.py | 2 +- openrag/api/routers/admin/monitoring.py | 2 +- openrag/api/routers/admin/partitions.py | 2 +- openrag/api/routers/admin/tools.py | 2 +- openrag/api/routers/admin/users.py | 2 +- openrag/api/routers/admin/workspaces.py | 2 +- openrag/api/routers/auth/oidc.py | 2 +- openrag/api/routers/user/chat.py | 4 +- openrag/api/routers/user/extract.py | 2 +- openrag/api/routers/user/search.py | 2 +- openrag/app_front.py | 2 +- openrag/components/indexer/chunker/chunker.py | 2 +- .../components/indexer/embeddings/openai.py | 8 +- .../components/indexer/loaders/__init__.py | 2 +- .../indexer/loaders/audio/local_whisper.py | 2 +- .../indexer/loaders/audio/openai.py | 2 +- openrag/components/indexer/loaders/base.py | 4 +- openrag/components/indexer/loaders/doc.py | 2 +- openrag/components/indexer/loaders/docx.py | 2 +- openrag/components/indexer/loaders/image.py | 2 +- .../indexer/loaders/pdf_loaders/docling.py | 2 +- .../indexer/loaders/pdf_loaders/docling2.py | 2 +- .../indexer/loaders/pdf_loaders/dotsocr.py | 2 +- .../indexer/loaders/pdf_loaders/marker.py | 2 +- .../indexer/loaders/pdf_loaders/openai.py | 2 +- .../indexer/loaders/pdf_loaders/pymupdf.py | 2 +- .../components/indexer/loaders/pptx_loader.py | 2 +- .../components/indexer/loaders/txt_loader.py | 2 +- openrag/components/llm.py | 2 +- openrag/components/reranker/infinity.py | 2 +- openrag/components/reranker/openai.py | 2 +- openrag/components/utils.py | 2 +- .../components/websearch/content_fetcher.py | 2 +- .../components/websearch/providers/staan.py | 2 +- openrag/components/websearch/service.py | 2 +- openrag/core/config/__init__.py | 37 +++++ .../observability}/monitoring.py | 0 .../logger.py => core/utils/logging.py} | 2 +- .../utils/test_logging.py} | 2 +- openrag/di/container.py | 2 +- openrag/scripts/backup.py | 2 +- openrag/scripts/restore.py | 2 +- openrag/services/auth/refresh.py | 2 +- .../services/inference/_circuit_breaker.py | 2 +- openrag/services/inference/_retry.py | 2 +- openrag/services/inference/healthcheck.py | 2 +- openrag/services/inference/ollama_client.py | 2 +- .../services/inference/reranker_clients.py | 2 +- openrag/services/inference/vllm_client.py | 2 +- .../services/orchestrators/auth_service.py | 2 +- .../orchestrators/conversion_service.py | 2 +- .../orchestrators/indexing_service.py | 2 +- openrag/services/orchestrators/mcp_service.py | 2 +- .../orchestrators/partition_service.py | 2 +- .../services/orchestrators/query_service.py | 2 +- .../orchestrators/retrieval_service.py | 2 +- .../services/orchestrators/user_service.py | 2 +- .../orchestrators/workspace_service.py | 2 +- openrag/services/persistence/connection.py | 2 +- .../1.add_created_at_temporal_fields.py | 2 +- .../persistence/migrations/milvus/migrate.py | 2 +- openrag/services/workers/batch_ingest.py | 2 +- openrag/services/workers/bootstrap.py | 2 +- .../workers/parsers/doc_serializer.py | 2 +- .../workers/parsers/docling_workers.py | 2 +- .../workers/parsers/marker_workers.py | 6 +- .../workers/parsers/whisper_workers.py | 6 +- openrag/services/workers/ray_utils.py | 2 +- openrag/test_auth_router.py | 6 +- .../utils/test_external_resource_errors.py | 129 ------------------ 80 files changed, 128 insertions(+), 216 deletions(-) rename openrag/{utils => core/observability}/monitoring.py (100%) rename openrag/{utils/logger.py => core/utils/logging.py} (98%) rename openrag/{utils/test_logger.py => core/utils/test_logging.py} (94%) delete mode 100644 openrag/utils/test_external_resource_errors.py diff --git a/CLAUDE.md b/CLAUDE.md index ab02288d7..b40c41591 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -312,7 +312,7 @@ All custom exceptions inherit from `OpenRAGError` (`openrag/utils/exceptions/`): Uses Loguru with structured logging: ```python -from utils.logger import get_logger +from core.utils.logging import get_logger logger = get_logger() logger.bind(file_id=file_id, partition=partition).info("Message") ``` @@ -323,7 +323,7 @@ Use absolute imports from the `openrag/` directory (which is the Python path roo ```python # Correct - absolute imports from components.ray_utils import call_ray_actor_with_timeout -from utils.logger import get_logger +from core.utils.logging import get_logger from config import load_config # Avoid relative imports across packages diff --git a/docs/content/docs/documentation/milvus_migration.mdx b/docs/content/docs/documentation/milvus_migration.mdx index cb362d4af..cf73e3ce1 100644 --- a/docs/content/docs/documentation/milvus_migration.mdx +++ b/docs/content/docs/documentation/milvus_migration.mdx @@ -276,7 +276,7 @@ Milvus migration: (schema version N-1 → N) """ from pymilvus import DataType, MilvusClient -from utils.logger import get_logger +from core.utils.logging import get_logger TARGET_VERSION = N # replace with the actual version number diff --git a/openrag/api/dependencies/auth.py b/openrag/api/dependencies/auth.py index 3d930ecba..9eca30850 100644 --- a/openrag/api/dependencies/auth.py +++ b/openrag/api/dependencies/auth.py @@ -1,9 +1,9 @@ import os from core.utils.exceptions import OpenRAGError +from core.utils.logging import get_logger from di.providers import get_auth_service, get_config, get_job_service, get_partition_service from fastapi import Depends, HTTPException, Request, status -from utils.logger import get_logger logger = get_logger() diff --git a/openrag/api/dependencies/llm.py b/openrag/api/dependencies/llm.py index 6334257a3..9b12630c1 100644 --- a/openrag/api/dependencies/llm.py +++ b/openrag/api/dependencies/llm.py @@ -1,10 +1,10 @@ import consts import openai from api.dependencies.auth import SUPER_ADMIN_MODE +from core.utils.logging import get_logger from di.providers import get_config from fastapi import Depends, HTTPException, status from openai import AsyncOpenAI -from utils.logger import get_logger logger = get_logger() diff --git a/openrag/api/error_handlers.py b/openrag/api/error_handlers.py index 294442dc5..952667e4a 100644 --- a/openrag/api/error_handlers.py +++ b/openrag/api/error_handlers.py @@ -38,9 +38,9 @@ StorageError, ValidationError, ) +from core.utils.logging import get_logger from fastapi import FastAPI, Request from fastapi.responses import JSONResponse -from utils.logger import get_logger logger = get_logger() diff --git a/openrag/api/main.py b/openrag/api/main.py index bf1b79638..89012cadf 100644 --- a/openrag/api/main.py +++ b/openrag/api/main.py @@ -52,6 +52,7 @@ from api.routers.user.health import router as health_router from api.routers.user.search import router as search_router from config import load_config +from core.utils.logging import get_logger from di.container import ServiceContainer from di.providers import set_container from di.workers import ensure_worker_bootstrap @@ -61,7 +62,6 @@ from fastapi.openapi.utils import get_openapi from fastapi.responses import JSONResponse, RedirectResponse from fastapi.staticfiles import StaticFiles -from utils.logger import get_logger # pydub 0.25.1 ships invalid-escape regex literals; the warning is upstream. warnings.filterwarnings("ignore", category=SyntaxWarning, module="pydub") diff --git a/openrag/api/mcp/server.py b/openrag/api/mcp/server.py index 5b924eb02..4d5ca5c11 100644 --- a/openrag/api/mcp/server.py +++ b/openrag/api/mcp/server.py @@ -37,13 +37,13 @@ ) from config import load_config from core.utils.log_tail import app_log_file +from core.utils.logging import get_logger from di.container import ServiceContainer from di.workers import ensure_worker_bootstrap from mcp.server.fastmcp import FastMCP from starlette.middleware.base import BaseHTTPMiddleware from starlette.requests import Request from starlette.responses import JSONResponse -from utils.logger import get_logger logger = get_logger() diff --git a/openrag/api/middleware/auth.py b/openrag/api/middleware/auth.py index de428d8d6..6a87f06b6 100644 --- a/openrag/api/middleware/auth.py +++ b/openrag/api/middleware/auth.py @@ -36,10 +36,10 @@ from components.auth.refresh import refresh_session_if_needed from core.config.auth import AuthBypassConfig +from core.utils.logging import get_logger from fastapi import Request from fastapi.responses import JSONResponse, RedirectResponse from starlette.middleware.base import BaseHTTPMiddleware -from utils.logger import get_logger logger = get_logger() diff --git a/openrag/api/middleware/instrumentation.py b/openrag/api/middleware/instrumentation.py index 9bb11ee0c..c800bb140 100644 --- a/openrag/api/middleware/instrumentation.py +++ b/openrag/api/middleware/instrumentation.py @@ -9,9 +9,9 @@ import time +from core.observability.monitoring import record_request from fastapi import Request from starlette.middleware.base import BaseHTTPMiddleware -from utils.monitoring import record_request # Paths to exclude from metric recording — avoids self-referential noise # and inflated counters on probe traffic. diff --git a/openrag/api/routers/admin/cluster.py b/openrag/api/routers/admin/cluster.py index 651a66cd7..db6f33db9 100644 --- a/openrag/api/routers/admin/cluster.py +++ b/openrag/api/routers/admin/cluster.py @@ -1,9 +1,9 @@ from api.dependencies.auth import require_admin +from core.utils.logging import get_logger from di.workers import list_ray_actors as list_ray_actor_states from di.workers import restart_ray_actor from fastapi import APIRouter, Depends, HTTPException, status from fastapi.responses import JSONResponse -from utils.logger import get_logger logger = get_logger() diff --git a/openrag/api/routers/admin/indexing.py b/openrag/api/routers/admin/indexing.py index 188ccd43e..e43a354af 100644 --- a/openrag/api/routers/admin/indexing.py +++ b/openrag/api/routers/admin/indexing.py @@ -30,6 +30,7 @@ from api.routers.admin.task_logs import collect_task_logs from components.indexer.utils.files import sanitize_filename, save_file_to_disk from core.utils.log_tail import app_log_file +from core.utils.logging import get_logger from di.providers import get_auth_service, get_config, get_indexing_service, get_partition_service from fastapi import ( APIRouter, @@ -42,7 +43,6 @@ status, ) from fastapi.responses import JSONResponse -from utils.logger import get_logger logger = get_logger() diff --git a/openrag/api/routers/admin/monitoring.py b/openrag/api/routers/admin/monitoring.py index 532344fda..401cab5bf 100644 --- a/openrag/api/routers/admin/monitoring.py +++ b/openrag/api/routers/admin/monitoring.py @@ -7,9 +7,9 @@ import asyncio from api.dependencies.auth import require_admin +from core.observability.monitoring import get_metrics from fastapi import APIRouter, Depends from fastapi.responses import Response -from utils.monitoring import get_metrics router = APIRouter() diff --git a/openrag/api/routers/admin/partitions.py b/openrag/api/routers/admin/partitions.py index 05301b39a..c95b42c2c 100644 --- a/openrag/api/routers/admin/partitions.py +++ b/openrag/api/routers/admin/partitions.py @@ -17,10 +17,10 @@ require_partition_owner, require_partition_viewer, ) +from core.utils.logging import get_logger from di.providers import get_partition_service from fastapi import APIRouter, Depends, Form, HTTPException, Request, Response, status from fastapi.responses import JSONResponse -from utils.logger import get_logger logger = get_logger() router = APIRouter() diff --git a/openrag/api/routers/admin/tools.py b/openrag/api/routers/admin/tools.py index e00222ea5..34598a18d 100644 --- a/openrag/api/routers/admin/tools.py +++ b/openrag/api/routers/admin/tools.py @@ -18,10 +18,10 @@ ) from api.schemas.admin.tools import ToolInfo from components.indexer.utils.files import save_file_to_disk +from core.utils.logging import get_logger from di.providers import get_config, get_conversion_service from fastapi import APIRouter, Depends, Form, HTTPException, UploadFile, status from fastapi.responses import JSONResponse -from utils.logger import get_logger logger = get_logger() diff --git a/openrag/api/routers/admin/users.py b/openrag/api/routers/admin/users.py index 7b1801e39..529a96066 100644 --- a/openrag/api/routers/admin/users.py +++ b/openrag/api/routers/admin/users.py @@ -16,10 +16,10 @@ from api.dependencies.auth import current_user, require_admin, require_admin_or_self from api.schemas.admin.users import UserCreate, UserPublic, UserUpdate +from core.utils.logging import get_logger from di.providers import get_user_service from fastapi import APIRouter, Depends, HTTPException, Response, status from fastapi.responses import JSONResponse -from utils.logger import get_logger logger = get_logger() router = APIRouter() diff --git a/openrag/api/routers/admin/workspaces.py b/openrag/api/routers/admin/workspaces.py index 646936b06..53ba7dba3 100644 --- a/openrag/api/routers/admin/workspaces.py +++ b/openrag/api/routers/admin/workspaces.py @@ -12,9 +12,9 @@ from api.dependencies.auth import require_partition_editor, require_partition_owner, require_partition_viewer from api.schemas.admin.workspaces import AddFilesRequest, CreateWorkspaceRequest +from core.utils.logging import get_logger from di.providers import get_workspace_service from fastapi import APIRouter, Depends, HTTPException, status -from utils.logger import get_logger router = APIRouter() logger = get_logger() diff --git a/openrag/api/routers/auth/oidc.py b/openrag/api/routers/auth/oidc.py index b022f2587..2df4ff4d1 100644 --- a/openrag/api/routers/auth/oidc.py +++ b/openrag/api/routers/auth/oidc.py @@ -25,10 +25,10 @@ from components.auth import StateCookieSerializer from core.utils.exceptions import OpenRAGError +from core.utils.logging import get_logger from di.providers import get_auth_service from fastapi import APIRouter, Depends, Form, HTTPException, Request, Response, status from fastapi.responses import JSONResponse, RedirectResponse -from utils.logger import get_logger logger = get_logger() router = APIRouter() diff --git a/openrag/api/routers/user/chat.py b/openrag/api/routers/user/chat.py index fd29157c8..3e57ba8f7 100644 --- a/openrag/api/routers/user/chat.py +++ b/openrag/api/routers/user/chat.py @@ -33,11 +33,11 @@ from components.indexer.utils.text_sanitizer import sanitize_text from components.utils import get_num_tokens from config import load_config +from core.utils.exceptions import OpenRAGError +from core.utils.logging import get_logger from di.providers import get_config, get_partition_service, get_query_service from fastapi import APIRouter, Body, Depends, HTTPException, Request, status from fastapi.responses import JSONResponse, StreamingResponse -from utils.exceptions.base import OpenRAGError -from utils.logger import get_logger logger = get_logger() router = APIRouter() diff --git a/openrag/api/routers/user/extract.py b/openrag/api/routers/user/extract.py index 49293e0be..4349a9ece 100644 --- a/openrag/api/routers/user/extract.py +++ b/openrag/api/routers/user/extract.py @@ -9,10 +9,10 @@ """ from api.dependencies.auth import current_user_or_admin_partitions_list +from core.utils.logging import get_logger from di.providers import get_conversion_service from fastapi import APIRouter, Depends, HTTPException, status from fastapi.responses import JSONResponse -from utils.logger import get_logger logger = get_logger() diff --git a/openrag/api/routers/user/search.py b/openrag/api/routers/user/search.py index 98d77f4d4..1455f81b1 100644 --- a/openrag/api/routers/user/search.py +++ b/openrag/api/routers/user/search.py @@ -16,10 +16,10 @@ require_partition_viewer, require_partitions_viewer, ) +from core.utils.logging import get_logger from di.providers import get_retrieval_service, get_workspace_service from fastapi import APIRouter, Depends, HTTPException, Query, Request, status from fastapi.responses import JSONResponse -from utils.logger import get_logger logger = get_logger() diff --git a/openrag/app_front.py b/openrag/app_front.py index bc3c757e2..607f837d7 100644 --- a/openrag/app_front.py +++ b/openrag/app_front.py @@ -9,9 +9,9 @@ from chainlit.config import config as cl_config from chainlit.context import get_context from consts import PARTITION_PREFIX +from core.utils.logging import get_logger from dotenv import load_dotenv from openai import AsyncOpenAI -from utils.logger import get_logger load_dotenv() logger = get_logger() diff --git a/openrag/components/indexer/chunker/chunker.py b/openrag/components/indexer/chunker/chunker.py index 90c1e238f..5427118bb 100644 --- a/openrag/components/indexer/chunker/chunker.py +++ b/openrag/components/indexer/chunker/chunker.py @@ -21,9 +21,9 @@ from core.models.chunk import Chunk as _CoreChunk from core.models.document import ProcessedDocument, TextBlock from core.prompts.contextualization_builder import wrap_chunk_with_context +from core.utils.logging import get_logger from langchain_core.documents.base import Document from langchain_core.messages import AIMessage, HumanMessage, SystemMessage -from utils.logger import get_logger if TYPE_CHECKING: from components.indexer.embeddings import BaseEmbedding diff --git a/openrag/components/indexer/embeddings/openai.py b/openrag/components/indexer/embeddings/openai.py index 2c23e6f90..1083b78c4 100644 --- a/openrag/components/indexer/embeddings/openai.py +++ b/openrag/components/indexer/embeddings/openai.py @@ -9,11 +9,15 @@ import httpx import openai from core.config.endpoints import EmbedderConfig +from core.utils.exceptions import ( + EmbeddingAPIError, + EmbeddingResponseError, + UnexpectedEmbeddingError, +) +from core.utils.logging import get_logger from langchain_core.documents.base import Document from openai import OpenAI from services.inference.vllm_client import VLLMEmbedder # noqa: F401 -from utils.exceptions.embeddings import * -from utils.logger import get_logger from .base import BaseEmbedding diff --git a/openrag/components/indexer/loaders/__init__.py b/openrag/components/indexer/loaders/__init__.py index 43fdafa90..7f2688e34 100644 --- a/openrag/components/indexer/loaders/__init__.py +++ b/openrag/components/indexer/loaders/__init__.py @@ -8,7 +8,7 @@ import pkgutil from pathlib import Path -from utils.logger import get_logger +from core.utils.logging import get_logger from .base import BaseLoader diff --git a/openrag/components/indexer/loaders/audio/local_whisper.py b/openrag/components/indexer/loaders/audio/local_whisper.py index ec1c70823..caa5cb3b1 100644 --- a/openrag/components/indexer/loaders/audio/local_whisper.py +++ b/openrag/components/indexer/loaders/audio/local_whisper.py @@ -23,6 +23,7 @@ from core.indexing.parsers.audio.local_whisper import LocalWhisperParser from core.models.document import Document as CoreDocument +from core.utils.logging import get_logger from langchain_core.documents.base import Document from services.workers.parsers.whisper_workers import ( # noqa: F401 (re-exported for legacy import paths) LocalWhisperLoader as _ServicesWhisperPool, @@ -31,7 +32,6 @@ WhisperActor, WhisperPool, ) -from utils.logger import get_logger from ..base import BaseLoader diff --git a/openrag/components/indexer/loaders/audio/openai.py b/openrag/components/indexer/loaders/audio/openai.py index 695e99528..4e8f976f0 100644 --- a/openrag/components/indexer/loaders/audio/openai.py +++ b/openrag/components/indexer/loaders/audio/openai.py @@ -19,10 +19,10 @@ from core.indexing.parsers.audio.client_based import ClientAudioParser from core.models.document import Document as CoreDocument from core.models.document import DocumentType +from core.utils.logging import get_logger from langchain_core.documents.base import Document from services.inference.parsers.openai_audio import OpenAIAudioClient from services.workers.parsers.whisper_workers import detect_language_via_actor -from utils.logger import get_logger from ..base import BaseLoader diff --git a/openrag/components/indexer/loaders/base.py b/openrag/components/indexer/loaders/base.py index ccc91ebac..8c9018263 100644 --- a/openrag/components/indexer/loaders/base.py +++ b/openrag/components/indexer/loaders/base.py @@ -19,13 +19,13 @@ ensure_png_compatible_mode, # noqa: F401 (re-exported for legacy import path) pil_to_png_bytes, ) +from core.utils.external_errors import is_external_resource_error +from core.utils.logging import get_logger from langchain_core.messages import HumanMessage from langchain_openai import ChatOpenAI from openai import BadRequestError from PIL import Image from tqdm.asyncio import tqdm -from utils.external_resource_errors import is_external_resource_error -from utils.logger import get_logger logger = get_logger() diff --git a/openrag/components/indexer/loaders/doc.py b/openrag/components/indexer/loaders/doc.py index 46b4e1cb4..cf0c7d882 100644 --- a/openrag/components/indexer/loaders/doc.py +++ b/openrag/components/indexer/loaders/doc.py @@ -17,9 +17,9 @@ from core.indexing.parsers.doc_parser import DocParser from core.models.document import Document as CoreDocument from core.models.document import DocumentType +from core.utils.logging import get_logger from langchain_core.documents.base import Document as LCDocument from PIL import Image -from utils.logger import get_logger from .base import BaseLoader diff --git a/openrag/components/indexer/loaders/docx.py b/openrag/components/indexer/loaders/docx.py index 39b48143d..7688f87bc 100644 --- a/openrag/components/indexer/loaders/docx.py +++ b/openrag/components/indexer/loaders/docx.py @@ -18,9 +18,9 @@ from core.indexing.parsers.docx_parser import DocxParser from core.models.document import Document as CoreDocument from core.models.document import DocumentType +from core.utils.logging import get_logger from langchain_core.documents.base import Document from PIL import Image -from utils.logger import get_logger from .base import BaseLoader, ensure_png_compatible_mode diff --git a/openrag/components/indexer/loaders/image.py b/openrag/components/indexer/loaders/image.py index 6f4280692..9aa418b8c 100644 --- a/openrag/components/indexer/loaders/image.py +++ b/openrag/components/indexer/loaders/image.py @@ -18,9 +18,9 @@ from core.models.document import Document as CoreDocument from core.models.document import DocumentType from core.utils.exceptions import OpenRAGError +from core.utils.logging import get_logger from langchain_core.documents import Document from PIL import Image -from utils.logger import get_logger from .base import BaseLoader diff --git a/openrag/components/indexer/loaders/pdf_loaders/docling.py b/openrag/components/indexer/loaders/pdf_loaders/docling.py index 0e745d4bc..5b5d7363d 100644 --- a/openrag/components/indexer/loaders/pdf_loaders/docling.py +++ b/openrag/components/indexer/loaders/pdf_loaders/docling.py @@ -2,6 +2,7 @@ import torch from components.utils import SingletonMeta +from core.utils.logging import get_logger from docling.backend.pypdfium2_backend import PyPdfiumDocumentBackend from docling.datamodel.base_models import InputFormat from docling.datamodel.document import ConversionResult @@ -15,7 +16,6 @@ from docling.document_converter import DocumentConverter, PdfFormatOption from docling_core.types.doc.document import PictureItem from langchain_core.documents.base import Document -from utils.logger import get_logger from ..base import BaseLoader diff --git a/openrag/components/indexer/loaders/pdf_loaders/docling2.py b/openrag/components/indexer/loaders/pdf_loaders/docling2.py index 22049a155..86a19b182 100644 --- a/openrag/components/indexer/loaders/pdf_loaders/docling2.py +++ b/openrag/components/indexer/loaders/pdf_loaders/docling2.py @@ -15,13 +15,13 @@ from __future__ import annotations +from core.utils.logging import get_logger from langchain_core.documents.base import Document from services.workers.parsers.docling_workers import ( # noqa: F401 (re-exported for legacy paths) DoclingLoader, DoclingPool, DoclingWorker, ) -from utils.logger import get_logger from ..base import BaseLoader diff --git a/openrag/components/indexer/loaders/pdf_loaders/dotsocr.py b/openrag/components/indexer/loaders/pdf_loaders/dotsocr.py index 696679ef0..c72c3638c 100644 --- a/openrag/components/indexer/loaders/pdf_loaders/dotsocr.py +++ b/openrag/components/indexer/loaders/pdf_loaders/dotsocr.py @@ -1,6 +1,6 @@ +from core.utils.logging import logger # assuming you have a shared logger instance from PIL import Image from tqdm.asyncio import tqdm -from utils.logger import logger # assuming you have a shared logger instance from .openai import OpenAILoader diff --git a/openrag/components/indexer/loaders/pdf_loaders/marker.py b/openrag/components/indexer/loaders/pdf_loaders/marker.py index fce7b3116..111db223e 100644 --- a/openrag/components/indexer/loaders/pdf_loaders/marker.py +++ b/openrag/components/indexer/loaders/pdf_loaders/marker.py @@ -24,6 +24,7 @@ from core.indexing.parsers.pdf.marker import MarkerParser from core.models.document import Document as CoreDocument from core.models.document import DocumentType +from core.utils.logging import get_logger from langchain_core.documents.base import Document from PIL import Image from services.workers.parsers.marker_workers import ( # noqa: F401 (re-exported for legacy import paths) @@ -33,7 +34,6 @@ MarkerPool, MarkerWorker, ) -from utils.logger import get_logger from ..base import BaseLoader diff --git a/openrag/components/indexer/loaders/pdf_loaders/openai.py b/openrag/components/indexer/loaders/pdf_loaders/openai.py index 9aec8ba24..3ad7f6782 100644 --- a/openrag/components/indexer/loaders/pdf_loaders/openai.py +++ b/openrag/components/indexer/loaders/pdf_loaders/openai.py @@ -7,10 +7,10 @@ from pathlib import Path import pypdfium2 as pdfium +from core.utils.logging import logger from langchain.schema import Document from langchain_openai import ChatOpenAI from PIL import Image -from utils.logger import logger from ..base import BaseLoader diff --git a/openrag/components/indexer/loaders/pdf_loaders/pymupdf.py b/openrag/components/indexer/loaders/pdf_loaders/pymupdf.py index c856c9523..c46dd0d89 100644 --- a/openrag/components/indexer/loaders/pdf_loaders/pymupdf.py +++ b/openrag/components/indexer/loaders/pdf_loaders/pymupdf.py @@ -17,9 +17,9 @@ from core.indexing.parsers.pdf.pymupdf import PyMuPDFParser from core.models.document import Document as CoreDocument from core.models.document import DocumentType +from core.utils.logging import get_logger from langchain_core.documents.base import Document from PIL import Image -from utils.logger import get_logger from ..base import BaseLoader diff --git a/openrag/components/indexer/loaders/pptx_loader.py b/openrag/components/indexer/loaders/pptx_loader.py index fbbb2a8d1..18b14f177 100644 --- a/openrag/components/indexer/loaders/pptx_loader.py +++ b/openrag/components/indexer/loaders/pptx_loader.py @@ -15,9 +15,9 @@ from core.indexing.parsers.pptx_parser import PptxParser from core.models.document import Document as CoreDocument from core.models.document import DocumentType +from core.utils.logging import get_logger from langchain_core.documents.base import Document from PIL import Image -from utils.logger import get_logger from .base import BaseLoader diff --git a/openrag/components/indexer/loaders/txt_loader.py b/openrag/components/indexer/loaders/txt_loader.py index 775b798c0..9e5154c10 100644 --- a/openrag/components/indexer/loaders/txt_loader.py +++ b/openrag/components/indexer/loaders/txt_loader.py @@ -19,8 +19,8 @@ from core.indexing.parsers.text_parser import TextParser from core.models.document import Document as CoreDocument from core.models.document import DocumentType +from core.utils.logging import get_logger from langchain_core.documents.base import Document -from utils.logger import get_logger logger = get_logger() diff --git a/openrag/components/llm.py b/openrag/components/llm.py index e4151879f..a2873941c 100644 --- a/openrag/components/llm.py +++ b/openrag/components/llm.py @@ -9,8 +9,8 @@ import httpx from config.models import LLMConfig +from core.utils.logging import get_logger from services.inference.vllm_client import VLLMClient # noqa: F401 -from utils.logger import get_logger logger = get_logger() diff --git a/openrag/components/reranker/infinity.py b/openrag/components/reranker/infinity.py index 32626466d..6135c4bec 100644 --- a/openrag/components/reranker/infinity.py +++ b/openrag/components/reranker/infinity.py @@ -5,12 +5,12 @@ import asyncio +from core.utils.logging import get_logger from infinity_client import Client from infinity_client.api.default import rerank from infinity_client.models import RerankInput, ReRankResult from langchain_core.documents.base import Document from services.inference.reranker_clients import InfinityReranker as InfinityRerankerAdapter # noqa: F401 -from utils.logger import get_logger from .base import BaseReranker diff --git a/openrag/components/reranker/openai.py b/openrag/components/reranker/openai.py index 7bca4fbcf..9af3c5947 100644 --- a/openrag/components/reranker/openai.py +++ b/openrag/components/reranker/openai.py @@ -6,9 +6,9 @@ import asyncio import httpx +from core.utils.logging import get_logger from langchain_core.documents.base import Document from services.inference.reranker_clients import OpenAIReranker as OpenAIRerankerAdapter # noqa: F401 -from utils.logger import get_logger from .base import BaseReranker diff --git a/openrag/components/utils.py b/openrag/components/utils.py index 16529f6af..4931ce226 100644 --- a/openrag/components/utils.py +++ b/openrag/components/utils.py @@ -6,13 +6,13 @@ from typing import ClassVar from config import load_config +from core.utils.logging import get_logger from fast_langdetect import LangDetectConfig, LangDetector from langchain_core.documents.base import Document from services.inference.distributed_semaphore import ( DistributedSemaphore, # noqa: F401 DistributedSemaphoreActor, # noqa: F401 ) -from utils.logger import get_logger SOURCE_SEPARATOR = "-" * 10 + "\n\n" diff --git a/openrag/components/websearch/content_fetcher.py b/openrag/components/websearch/content_fetcher.py index 27fa9d033..d34dd927f 100644 --- a/openrag/components/websearch/content_fetcher.py +++ b/openrag/components/websearch/content_fetcher.py @@ -5,8 +5,8 @@ import httpx import lxml.html from components.websearch.base import WebResult +from core.utils.logging import get_logger from html_to_markdown import convert -from utils.logger import get_logger logger = get_logger() diff --git a/openrag/components/websearch/providers/staan.py b/openrag/components/websearch/providers/staan.py index 6fedaf3ee..a8f0a300b 100644 --- a/openrag/components/websearch/providers/staan.py +++ b/openrag/components/websearch/providers/staan.py @@ -1,6 +1,6 @@ import httpx from components.websearch.base import BaseWebSearchProvider, WebResult -from utils.logger import get_logger +from core.utils.logging import get_logger logger = get_logger() diff --git a/openrag/components/websearch/service.py b/openrag/components/websearch/service.py index 8ef0e93bd..40f5bdd14 100644 --- a/openrag/components/websearch/service.py +++ b/openrag/components/websearch/service.py @@ -1,6 +1,6 @@ from components.websearch.base import BaseWebSearchProvider, WebResult from components.websearch.content_fetcher import ContentFetcher -from utils.logger import get_logger +from core.utils.logging import get_logger logger = get_logger() diff --git a/openrag/core/config/__init__.py b/openrag/core/config/__init__.py index e69de29bb..7f91fb315 100644 --- a/openrag/core/config/__init__.py +++ b/openrag/core/config/__init__.py @@ -0,0 +1,37 @@ +"""OpenRAG configuration package. + +Public API: + load_config() — load config (cached singleton, or fresh with overrides) + Settings — root Pydantic model + get_settings() — cached singleton accessor +""" + +from functools import lru_cache + +from .root import Settings + + +@lru_cache +def get_settings() -> Settings: + """Cached singleton — one Settings instance per process.""" + from .loader import load_config as _load + + return _load() + + +def load_config(config_path=None, overrides=None) -> Settings: + """Return the cached Pydantic Settings singleton. + + The ``config_path`` parameter is kept for backward compatibility. + Use ``OPENRAG_CONF_DIR`` env var to override the config directory. + + The ``overrides`` parameter bypasses the cache (useful for tests). + """ + if overrides or config_path: + from .loader import load_config as _load + + return _load(conf_dir=config_path, overrides=overrides) + return get_settings() + + +__all__ = ["load_config", "Settings", "get_settings"] diff --git a/openrag/utils/monitoring.py b/openrag/core/observability/monitoring.py similarity index 100% rename from openrag/utils/monitoring.py rename to openrag/core/observability/monitoring.py diff --git a/openrag/utils/logger.py b/openrag/core/utils/logging.py similarity index 98% rename from openrag/utils/logger.py rename to openrag/core/utils/logging.py index a43031636..9cd1b8b73 100644 --- a/openrag/utils/logger.py +++ b/openrag/core/utils/logging.py @@ -1,7 +1,7 @@ import os import sys -from config import load_config +from core.config import load_config from core.utils.log_tail import app_log_file from loguru import logger diff --git a/openrag/utils/test_logger.py b/openrag/core/utils/test_logging.py similarity index 94% rename from openrag/utils/test_logger.py rename to openrag/core/utils/test_logging.py index 8b6f78cef..7d742d0c4 100644 --- a/openrag/utils/test_logger.py +++ b/openrag/core/utils/test_logging.py @@ -1,4 +1,4 @@ -from utils.logger import escape_markup, mask_email +from core.utils.logging import escape_markup, mask_email def test_mask_email_keeps_first_char_and_domain(): diff --git a/openrag/di/container.py b/openrag/di/container.py index 3aa3022f8..66a4a46a0 100644 --- a/openrag/di/container.py +++ b/openrag/di/container.py @@ -26,6 +26,7 @@ from core.embeddings import embedder_registry from core.llm import llm_registry from core.rerankers import reranker_registry +from core.utils.logging import get_logger from core.vlm import vlm_registry from di.embedders import register_embedders from di.llms import register_llms @@ -33,7 +34,6 @@ from di.rerankers import register_rerankers from di.vector_stores import create_vector_store from di.vlms import register_vlms -from utils.logger import get_logger if TYPE_CHECKING: from core.config.root import Settings diff --git a/openrag/scripts/backup.py b/openrag/scripts/backup.py index cea613eb3..020ec4fd1 100644 --- a/openrag/scripts/backup.py +++ b/openrag/scripts/backup.py @@ -5,11 +5,11 @@ import sys from typing import IO, Any +from core.utils.logging import get_logger from pymilvus import Collection, connections from services.persistence.schema import files as files_table from services.persistence.schema import partitions as partitions_table from sqlalchemy import create_engine, select -from utils.logger import get_logger def _list_partitions(conn) -> list[dict]: diff --git a/openrag/scripts/restore.py b/openrag/scripts/restore.py index 1e13ff889..75ebcdb32 100644 --- a/openrag/scripts/restore.py +++ b/openrag/scripts/restore.py @@ -5,6 +5,7 @@ import time from typing import IO, Any +from core.utils.logging import get_logger from pymilvus import MilvusClient from services.persistence.schema import files as files_table from services.persistence.schema import partition_memberships @@ -12,7 +13,6 @@ from services.persistence.schema import users as users_table from sqlalchemy import create_engine, select from sqlalchemy.dialects.postgresql import insert as pg_insert -from utils.logger import get_logger def _list_partitions(conn) -> list[dict]: diff --git a/openrag/services/auth/refresh.py b/openrag/services/auth/refresh.py index 2f72dfec5..77cf137aa 100644 --- a/openrag/services/auth/refresh.py +++ b/openrag/services/auth/refresh.py @@ -40,9 +40,9 @@ from datetime import datetime, timedelta from typing import Any +from core.utils.logging import get_logger from services.auth.deps import get_oidc_client from services.auth.session_tokens import decrypt_token, encrypt_token -from utils.logger import get_logger _REFRESH_BUFFER = timedelta(seconds=60) _STAMPEDE_WINDOW = timedelta(seconds=5) diff --git a/openrag/services/inference/_circuit_breaker.py b/openrag/services/inference/_circuit_breaker.py index 2d332699c..79b1afc97 100644 --- a/openrag/services/inference/_circuit_breaker.py +++ b/openrag/services/inference/_circuit_breaker.py @@ -4,8 +4,8 @@ import httpx from aiobreaker import CircuitBreaker, CircuitBreakerError, CircuitBreakerListener from core.utils.exceptions import InferenceConnectionError, LLMParsingError, OpenRAGError +from core.utils.logging import get_logger from prometheus_client import Gauge -from utils.logger import get_logger logger = get_logger() diff --git a/openrag/services/inference/_retry.py b/openrag/services/inference/_retry.py index f74ff6bc1..734945ea4 100644 --- a/openrag/services/inference/_retry.py +++ b/openrag/services/inference/_retry.py @@ -1,5 +1,6 @@ import httpx from core.utils.exceptions import OpenRAGError +from core.utils.logging import get_logger from tenacity import ( RetryCallState, retry, @@ -7,7 +8,6 @@ stop_after_attempt, wait_exponential_jitter, ) -from utils.logger import get_logger logger = get_logger() diff --git a/openrag/services/inference/healthcheck.py b/openrag/services/inference/healthcheck.py index 7cf259006..2e3d4ef02 100644 --- a/openrag/services/inference/healthcheck.py +++ b/openrag/services/inference/healthcheck.py @@ -11,7 +11,7 @@ from enum import Enum import httpx -from utils.logger import get_logger +from core.utils.logging import get_logger logger = get_logger() diff --git a/openrag/services/inference/ollama_client.py b/openrag/services/inference/ollama_client.py index b200546b5..19259f760 100644 --- a/openrag/services/inference/ollama_client.py +++ b/openrag/services/inference/ollama_client.py @@ -22,7 +22,7 @@ InferenceError, InferenceTimeoutError, ) -from utils.logger import get_logger +from core.utils.logging import get_logger from ._circuit_breaker import with_circuit_breaker from ._retry import with_retry diff --git a/openrag/services/inference/reranker_clients.py b/openrag/services/inference/reranker_clients.py index 2e6366f3e..15fbf10e6 100644 --- a/openrag/services/inference/reranker_clients.py +++ b/openrag/services/inference/reranker_clients.py @@ -11,7 +11,7 @@ import httpx from core.rerankers import Reranker, reranker_registry from core.utils.exceptions import InferenceConnectionError, InferenceTimeoutError -from utils.logger import get_logger +from core.utils.logging import get_logger from ._circuit_breaker import with_circuit_breaker from ._retry import with_retry diff --git a/openrag/services/inference/vllm_client.py b/openrag/services/inference/vllm_client.py index 1b6fd9e4b..80f56f6ac 100644 --- a/openrag/services/inference/vllm_client.py +++ b/openrag/services/inference/vllm_client.py @@ -26,8 +26,8 @@ InferenceError, InferenceTimeoutError, ) +from core.utils.logging import get_logger from core.vlm import VLM, vlm_registry -from utils.logger import get_logger from ._circuit_breaker import with_circuit_breaker from ._retry import with_retry diff --git a/openrag/services/orchestrators/auth_service.py b/openrag/services/orchestrators/auth_service.py index 78b98c8f9..f14a8d4a6 100644 --- a/openrag/services/orchestrators/auth_service.py +++ b/openrag/services/orchestrators/auth_service.py @@ -36,7 +36,7 @@ ) from core.models.user import OIDCSession, User from core.utils.exceptions import AuthError, OpenRAGError -from utils.logger import get_logger, mask_email +from core.utils.logging import get_logger, mask_email if TYPE_CHECKING: from core.config.auth import OIDCConfig diff --git a/openrag/services/orchestrators/conversion_service.py b/openrag/services/orchestrators/conversion_service.py index 470b8b1d1..a77631f41 100644 --- a/openrag/services/orchestrators/conversion_service.py +++ b/openrag/services/orchestrators/conversion_service.py @@ -28,8 +28,8 @@ from typing import TYPE_CHECKING +from core.utils.logging import get_logger from core.utils.text import sanitize_extracted_text -from utils.logger import get_logger if TYPE_CHECKING: from core.indexing.serializer import FileSerializer diff --git a/openrag/services/orchestrators/indexing_service.py b/openrag/services/orchestrators/indexing_service.py index a98bb9997..9095efab1 100644 --- a/openrag/services/orchestrators/indexing_service.py +++ b/openrag/services/orchestrators/indexing_service.py @@ -17,7 +17,7 @@ from typing import TYPE_CHECKING from components.indexer.utils.files import extract_temporal_fields -from utils.logger import get_logger +from core.utils.logging import get_logger if TYPE_CHECKING: from core.indexing.dispatcher import IndexingDispatcher diff --git a/openrag/services/orchestrators/mcp_service.py b/openrag/services/orchestrators/mcp_service.py index ef6a21b31..566cd3c0a 100644 --- a/openrag/services/orchestrators/mcp_service.py +++ b/openrag/services/orchestrators/mcp_service.py @@ -34,8 +34,8 @@ import httpx from core.utils.log_tail import collect_task_logs +from core.utils.logging import get_logger from core.utils.url_safety import is_blocked_address, is_safe_url -from utils.logger import get_logger if TYPE_CHECKING: from core.vector_stores import VectorStore diff --git a/openrag/services/orchestrators/partition_service.py b/openrag/services/orchestrators/partition_service.py index 6e5c4441b..ef882d0ff 100644 --- a/openrag/services/orchestrators/partition_service.py +++ b/openrag/services/orchestrators/partition_service.py @@ -34,7 +34,7 @@ UserNotFoundError, ValidationError, ) -from utils.logger import get_logger +from core.utils.logging import get_logger if TYPE_CHECKING: from core.ports.document_repo import DocumentRepository diff --git a/openrag/services/orchestrators/query_service.py b/openrag/services/orchestrators/query_service.py index fda3495a5..20f81e4c0 100644 --- a/openrag/services/orchestrators/query_service.py +++ b/openrag/services/orchestrators/query_service.py @@ -58,7 +58,7 @@ stream_with_source_filtering, ) from core.models.query import Query, SearchQueries -from utils.logger import get_logger +from core.utils.logging import get_logger if TYPE_CHECKING: from core.config.root import Settings diff --git a/openrag/services/orchestrators/retrieval_service.py b/openrag/services/orchestrators/retrieval_service.py index 43d509c5f..bf6487d12 100644 --- a/openrag/services/orchestrators/retrieval_service.py +++ b/openrag/services/orchestrators/retrieval_service.py @@ -35,7 +35,7 @@ _expand_with_related_chunks, ) from core.retrieval.rrf import rrf_reranking -from utils.logger import get_logger +from core.utils.logging import get_logger if TYPE_CHECKING: from core.config.root import Settings diff --git a/openrag/services/orchestrators/user_service.py b/openrag/services/orchestrators/user_service.py index e21211701..1848e05db 100644 --- a/openrag/services/orchestrators/user_service.py +++ b/openrag/services/orchestrators/user_service.py @@ -25,7 +25,7 @@ from typing import TYPE_CHECKING, Any from core.utils.exceptions import UserNotFoundError, ValidationError -from utils.logger import get_logger +from core.utils.logging import get_logger if TYPE_CHECKING: from core.models.user import UserCreate, UserUpdate diff --git a/openrag/services/orchestrators/workspace_service.py b/openrag/services/orchestrators/workspace_service.py index c77457f10..156a1f57a 100644 --- a/openrag/services/orchestrators/workspace_service.py +++ b/openrag/services/orchestrators/workspace_service.py @@ -23,7 +23,7 @@ import asyncio from typing import TYPE_CHECKING -from utils.logger import get_logger +from core.utils.logging import get_logger if TYPE_CHECKING: from core.ports.document_repo import DocumentRepository diff --git a/openrag/services/persistence/connection.py b/openrag/services/persistence/connection.py index 514085624..939607ae6 100644 --- a/openrag/services/persistence/connection.py +++ b/openrag/services/persistence/connection.py @@ -23,7 +23,7 @@ from typing import TYPE_CHECKING import asyncpg -from utils.logger import get_logger +from core.utils.logging import get_logger if TYPE_CHECKING: from core.config.infrastructure import RDBConfig diff --git a/openrag/services/persistence/migrations/milvus/1.add_created_at_temporal_fields.py b/openrag/services/persistence/migrations/milvus/1.add_created_at_temporal_fields.py index df22de6fe..e0651414c 100644 --- a/openrag/services/persistence/migrations/milvus/1.add_created_at_temporal_fields.py +++ b/openrag/services/persistence/migrations/milvus/1.add_created_at_temporal_fields.py @@ -33,9 +33,9 @@ import sys from config import load_config +from core.utils.logging import get_logger from pymilvus import DataType, MilvusClient from services.storage.milvus_store import SCHEMA_VERSION_PROPERTY_KEY -from utils.logger import get_logger TARGET_VERSION = 1 # The schema version this migration brings the collection to. diff --git a/openrag/services/persistence/migrations/milvus/migrate.py b/openrag/services/persistence/migrations/milvus/migrate.py index c732d1833..7d3cd8f13 100644 --- a/openrag/services/persistence/migrations/milvus/migrate.py +++ b/openrag/services/persistence/migrations/milvus/migrate.py @@ -37,9 +37,9 @@ from types import ModuleType from config import load_config +from core.utils.logging import get_logger from pymilvus import MilvusClient from services.storage.milvus_store import SCHEMA_VERSION_PROPERTY_KEY -from utils.logger import get_logger logger = get_logger() diff --git a/openrag/services/workers/batch_ingest.py b/openrag/services/workers/batch_ingest.py index 7a18d0c16..57ab7e9bd 100644 --- a/openrag/services/workers/batch_ingest.py +++ b/openrag/services/workers/batch_ingest.py @@ -4,8 +4,8 @@ from collections.abc import MutableMapping, Sequence from typing import Any +from core.utils.logging import get_logger from services.workers.pipeline_builder import IndexingPipeline -from utils.logger import get_logger _logger = get_logger().bind(component="batch_ingest") diff --git a/openrag/services/workers/bootstrap.py b/openrag/services/workers/bootstrap.py index e9c12326b..834d5dde3 100644 --- a/openrag/services/workers/bootstrap.py +++ b/openrag/services/workers/bootstrap.py @@ -22,7 +22,7 @@ from typing import TYPE_CHECKING import ray -from utils.logger import get_logger +from core.utils.logging import get_logger if TYPE_CHECKING: from core.config.root import Settings diff --git a/openrag/services/workers/parsers/doc_serializer.py b/openrag/services/workers/parsers/doc_serializer.py index 80c7d9e68..736d11a48 100644 --- a/openrag/services/workers/parsers/doc_serializer.py +++ b/openrag/services/workers/parsers/doc_serializer.py @@ -19,7 +19,7 @@ class DocSerializer: def __init__(self, data_dir=None, **kwargs) -> None: from config import load_config - from utils.logger import get_logger + from core.utils.logging import get_logger self.logger = get_logger() self.config = load_config() diff --git a/openrag/services/workers/parsers/docling_workers.py b/openrag/services/workers/parsers/docling_workers.py index 0a52e2ba2..9952777bf 100644 --- a/openrag/services/workers/parsers/docling_workers.py +++ b/openrag/services/workers/parsers/docling_workers.py @@ -18,6 +18,7 @@ from core.indexing.image_preprocessor import pil_to_png_bytes from core.indexing.parsers.document_parser import BasePooledParser from core.models.document import Document, DocumentType, ImageBlock, ProcessedDocument, TextBlock +from core.utils.logging import get_logger from docling.backend.pypdfium2_backend import PyPdfiumDocumentBackend from docling.datamodel.base_models import InputFormat from docling.datamodel.document import ConversionResult @@ -29,7 +30,6 @@ TableStructureOptions, ) from docling.document_converter import DocumentConverter, PdfFormatOption -from utils.logger import get_logger from ..ray_utils import call_ray_actor_with_timeout, retry_with_backoff diff --git a/openrag/services/workers/parsers/marker_workers.py b/openrag/services/workers/parsers/marker_workers.py index 36417bfb5..382cf6ec9 100644 --- a/openrag/services/workers/parsers/marker_workers.py +++ b/openrag/services/workers/parsers/marker_workers.py @@ -16,8 +16,8 @@ ProcessedDocument, TextBlock, ) +from core.utils.logging import get_logger from marker.converters.pdf import PdfConverter -from utils.logger import get_logger from ..ray_utils import call_ray_actor_with_timeout, retry_with_backoff @@ -34,7 +34,7 @@ def __init__(self): import os from config import load_config - from utils.logger import get_logger + from core.utils.logging import get_logger self.logger = get_logger() self.config = load_config() @@ -177,7 +177,7 @@ def __del__(self): class MarkerPool: def __init__(self): from config import load_config - from utils.logger import get_logger + from core.utils.logging import get_logger self.logger = get_logger() self.config = load_config() diff --git a/openrag/services/workers/parsers/whisper_workers.py b/openrag/services/workers/parsers/whisper_workers.py index 697f2fee2..ff2a959bf 100644 --- a/openrag/services/workers/parsers/whisper_workers.py +++ b/openrag/services/workers/parsers/whisper_workers.py @@ -11,8 +11,8 @@ ProcessedDocument, TextBlock, ) +from core.utils.logging import get_logger from faster_whisper import WhisperModel -from utils.logger import get_logger from ..ray_utils import call_ray_actor_with_timeout, retry_with_backoff @@ -40,7 +40,7 @@ class WhisperActor: def __init__(self): import torch from config import load_config - from utils.logger import get_logger + from core.utils.logging import get_logger self.logger = get_logger() self.config = load_config() @@ -103,7 +103,7 @@ class WhisperPool: def __init__(self): from config import load_config - from utils.logger import get_logger + from core.utils.logging import get_logger self.logger = get_logger() self.config = load_config() diff --git a/openrag/services/workers/ray_utils.py b/openrag/services/workers/ray_utils.py index 5061bc244..bcab7d81f 100644 --- a/openrag/services/workers/ray_utils.py +++ b/openrag/services/workers/ray_utils.py @@ -27,8 +27,8 @@ from typing import Any import ray +from core.utils.logging import get_logger from ray.exceptions import RayTaskError, TaskCancelledError -from utils.logger import get_logger logger = get_logger() diff --git a/openrag/test_auth_router.py b/openrag/test_auth_router.py index 73bb62d69..146788caa 100644 --- a/openrag/test_auth_router.py +++ b/openrag/test_auth_router.py @@ -34,7 +34,7 @@ # --------------------------------------------------------------------------- -_STUBBED_MODULES = ("utils", "utils.logger", "services.workers.bootstrap") +_STUBBED_MODULES = ("core.utils.logging", "services.workers.bootstrap") def _install_dependencies_stub() -> dict[str, types.ModuleType | None]: @@ -58,7 +58,7 @@ def _logger(): logger.bind = lambda *args, **kwargs: logger return logger - logger_stub = types.ModuleType("utils.logger") + logger_stub = types.ModuleType("core.utils.logging") logger_stub.escape_markup = lambda s: s.replace("\\", "\\\\").replace("<", "\\<").replace(">", "\\>") logger_stub.mask_email = ( lambda email: f"{email.partition('@')[0][0]}***@{email.partition('@')[2]}" @@ -66,7 +66,7 @@ def _logger(): else "***" ) logger_stub.get_logger = _logger - sys.modules["utils.logger"] = logger_stub + sys.modules["core.utils.logging"] = logger_stub return previous_modules diff --git a/openrag/utils/test_external_resource_errors.py b/openrag/utils/test_external_resource_errors.py deleted file mode 100644 index a85d6c025..000000000 --- a/openrag/utils/test_external_resource_errors.py +++ /dev/null @@ -1,129 +0,0 @@ -""" -Tests for external resource error detection utilities. -Related to: https://github.com/linagora/openrag/issues/182 -""" - -import pytest -from utils.external_resource_errors import is_external_resource_error - - -class TestIsExternalResourceError: - """Test suite for is_external_resource_error function.""" - - @pytest.mark.parametrize( - "error_msg,expected_code,url_contains", - [ - # Issue #182 - ( - "aiohttp.client_exceptions.ClientResponseError: 403, message='Forbidden', " - "url='https://upload.wikimedia.org/wikipedia/commons/thumb/d/d5/Logo.png'", - "403", - "upload.wikimedia.org", - ), - # Other HTTP status codes - ( - "ClientResponseError: 404, url='https://example.com/missing.png'", - "404", - "example.com", - ), - ( - "HTTPError: 401 Unauthorized for url: https://api.example.com/image.jpg", - "401", - "api.example.com", - ), - ( - "ClientResponseError: 429 Too Many Requests - https://cdn.example.com/img.png", - "429", - "cdn.example.com", - ), - # 5xx gateway errors - ( - "502 Bad Gateway: https://api.example.com/image.png", - "502", - "api.example.com", - ), - ( - "ClientResponseError: 503 Service Unavailable - https://cdn.example.com/img.png", - "503", - "cdn.example.com", - ), - # vLLM wrapped error (the real-world scenario) - ( - "openai.InternalServerError: Error code: 500 - {'error': {'message': " - "'litellm.InternalServerError: aiohttp.client_exceptions.ClientResponseError: " - "403, message=Forbidden, url=https://example.com/path/to/image.png'}}", - "403", - "example.com/path/to/image.png", - ), - ], - ) - def test_detects_http_errors_with_urls(self, error_msg, expected_code, url_contains): - """Test detection of HTTP errors with URL extraction.""" - is_external, status_code, url = is_external_resource_error(Exception(error_msg)) - - assert is_external is True - assert status_code == expected_code - assert url_contains in url - - @pytest.mark.parametrize( - "error_msg", - [ - "TimeoutError: Connection timed out while fetching resource", - "SSLError: Certificate verification failed", - "ConnectionError: Failed to connect to server", - "aiohttp.client_exceptions.ClientResponseError: some error", - "requests.exceptions.HTTPError: 500 Server Error", - ], - ) - def test_detects_error_indicators(self, error_msg): - """Test detection via error type indicators.""" - is_external, _, _ = is_external_resource_error(Exception(error_msg)) - assert is_external is True - - @pytest.mark.parametrize( - "error", - [ - Exception("ValueError: Invalid input parameter"), - Exception("Something went wrong during processing"), - TypeError("'NoneType' object is not subscriptable"), - AttributeError("'dict' object has no attribute 'content'"), - Exception(""), - # vLLM error without external cause details - Exception( - "openai.InternalServerError: Error code: 500 - {'error': {'message': " - "'litellm.InternalServerError: InternalServerError: OpenAIException'}}" - ), - ], - ) - def test_does_not_flag_internal_errors(self, error): - """Test that internal/generic errors are not flagged as external.""" - is_external, status_code, url = is_external_resource_error(error) - - assert is_external is False - assert status_code == "" - assert url == "" - - def test_extracts_url_with_query_params(self): - """Test URL extraction with query parameters.""" - error = Exception("403 Forbidden: https://api.example.com/image?id=123&size=large") - _, _, url = is_external_resource_error(error) - - assert "api.example.com/image?id=123" in url - - def test_indicator_substring_causes_false_positive(self): - """Document known limitation: indicator substrings cause false positives. - - This test documents that internal errors mentioning HTTP error class names - will be incorrectly classified as external. This is accepted because: - 1. Real error messages use these as exception class names, not prose - 2. Stricter matching (word boundaries) would break legitimate matches - like 'aiohttp.client_exceptions.ClientResponseError' - 3. This scenario is unlikely in practice - """ - error = Exception("InternalServerError: Failed to handle ClientResponseError in retry logic") - is_external, status_code, url = is_external_resource_error(error) - - # This IS classified as external (false positive) due to substring match - assert is_external is True - assert status_code == "" # No HTTP status code - assert url == "" # No URL From 53131f7c53635787c7a1e859cac23a52db145846 Mon Sep 17 00:00:00 2001 From: hedhoud <74668966+hedhoud@users.noreply.github.com> Date: Fri, 29 May 2026 14:28:10 +0200 Subject: [PATCH 02/11] refactor: move websearch to services --- openrag/core/config/retrieval.py | 2 +- openrag/di/container.py | 4 ++-- openrag/{components => services}/websearch/__init__.py | 0 openrag/{components => services}/websearch/base.py | 0 .../{components => services}/websearch/content_fetcher.py | 2 +- .../websearch/providers/__init__.py | 0 .../{components => services}/websearch/providers/staan.py | 2 +- openrag/{components => services}/websearch/service.py | 4 ++-- .../websearch/test_content_fetcher.py | 6 +++--- 9 files changed, 10 insertions(+), 10 deletions(-) rename openrag/{components => services}/websearch/__init__.py (100%) rename openrag/{components => services}/websearch/base.py (100%) rename openrag/{components => services}/websearch/content_fetcher.py (99%) rename openrag/{components => services}/websearch/providers/__init__.py (100%) rename openrag/{components => services}/websearch/providers/staan.py (94%) rename openrag/{components => services}/websearch/service.py (89%) rename openrag/{components => services}/websearch/test_content_fetcher.py (98%) diff --git a/openrag/core/config/retrieval.py b/openrag/core/config/retrieval.py index ca26abd57..75892e35a 100644 --- a/openrag/core/config/retrieval.py +++ b/openrag/core/config/retrieval.py @@ -122,7 +122,7 @@ class _BaseWebSearchConfig(ConfigMixin): fetch_max_results: int = 3 fetch_timeout: float = 1.0 fetch_max_tokens: int = 500 - fetch_verify_ssl: bool = False + fetch_verify_ssl: bool = True class StaanWebSearchConfig(_BaseWebSearchConfig): diff --git a/openrag/di/container.py b/openrag/di/container.py index 3aa3022f8..5dec0e9be 100644 --- a/openrag/di/container.py +++ b/openrag/di/container.py @@ -391,12 +391,12 @@ def query_service(self) -> QueryService: Shares the same core LLM construction as ``retrieval_service`` (built from ``settings.llm``); the web-search service comes from - the legacy ``WebSearchFactory`` (provider is ``None`` when + the ``WebSearchFactory`` (provider is ``None`` when ``WEBSEARCH_API_TOKEN`` is unset — web search silently disabled). """ if self._query_service is None: - from components.websearch import WebSearchFactory from services.orchestrators.query_service import QueryService + from services.websearch import WebSearchFactory settings = self._require_settings() llm_cfg = settings.llm.model_dump() diff --git a/openrag/components/websearch/__init__.py b/openrag/services/websearch/__init__.py similarity index 100% rename from openrag/components/websearch/__init__.py rename to openrag/services/websearch/__init__.py diff --git a/openrag/components/websearch/base.py b/openrag/services/websearch/base.py similarity index 100% rename from openrag/components/websearch/base.py rename to openrag/services/websearch/base.py diff --git a/openrag/components/websearch/content_fetcher.py b/openrag/services/websearch/content_fetcher.py similarity index 99% rename from openrag/components/websearch/content_fetcher.py rename to openrag/services/websearch/content_fetcher.py index 27fa9d033..59840fd34 100644 --- a/openrag/components/websearch/content_fetcher.py +++ b/openrag/services/websearch/content_fetcher.py @@ -4,8 +4,8 @@ import httpx import lxml.html -from components.websearch.base import WebResult from html_to_markdown import convert +from services.websearch.base import WebResult from utils.logger import get_logger logger = get_logger() diff --git a/openrag/components/websearch/providers/__init__.py b/openrag/services/websearch/providers/__init__.py similarity index 100% rename from openrag/components/websearch/providers/__init__.py rename to openrag/services/websearch/providers/__init__.py diff --git a/openrag/components/websearch/providers/staan.py b/openrag/services/websearch/providers/staan.py similarity index 94% rename from openrag/components/websearch/providers/staan.py rename to openrag/services/websearch/providers/staan.py index 6fedaf3ee..7f4ba755f 100644 --- a/openrag/components/websearch/providers/staan.py +++ b/openrag/services/websearch/providers/staan.py @@ -1,5 +1,5 @@ import httpx -from components.websearch.base import BaseWebSearchProvider, WebResult +from services.websearch.base import BaseWebSearchProvider, WebResult from utils.logger import get_logger logger = get_logger() diff --git a/openrag/components/websearch/service.py b/openrag/services/websearch/service.py similarity index 89% rename from openrag/components/websearch/service.py rename to openrag/services/websearch/service.py index 8ef0e93bd..2f01217ec 100644 --- a/openrag/components/websearch/service.py +++ b/openrag/services/websearch/service.py @@ -1,5 +1,5 @@ -from components.websearch.base import BaseWebSearchProvider, WebResult -from components.websearch.content_fetcher import ContentFetcher +from services.websearch.base import BaseWebSearchProvider, WebResult +from services.websearch.content_fetcher import ContentFetcher from utils.logger import get_logger logger = get_logger() diff --git a/openrag/components/websearch/test_content_fetcher.py b/openrag/services/websearch/test_content_fetcher.py similarity index 98% rename from openrag/components/websearch/test_content_fetcher.py rename to openrag/services/websearch/test_content_fetcher.py index 5ab74f81a..79d6af9ce 100644 --- a/openrag/components/websearch/test_content_fetcher.py +++ b/openrag/services/websearch/test_content_fetcher.py @@ -2,8 +2,8 @@ import httpx import pytest -from components.websearch.base import WebResult -from components.websearch.content_fetcher import ContentFetcher, _is_safe_url +from services.websearch.base import WebResult +from services.websearch.content_fetcher import ContentFetcher, _is_safe_url @pytest.fixture @@ -300,7 +300,7 @@ def test_fetch_verify_ssl_default_is_true(): ``fetch_verify_ssl`` defaults to ``True`` so server-side fetches of web-search results verify TLS unless an operator opts out. """ - from config.models import StaanWebSearchConfig + from core.config.retrieval import StaanWebSearchConfig cfg = StaanWebSearchConfig() assert cfg.fetch_verify_ssl is True From 5b6c2f70809b922e91bb185d55a9361e4d0e2254 Mon Sep 17 00:00:00 2001 From: hedhoud <74668966+hedhoud@users.noreply.github.com> Date: Fri, 29 May 2026 14:30:57 +0200 Subject: [PATCH 03/11] fix: address logger migration review comments --- openrag/app_front.py | 4 ++-- .../components/indexer/loaders/pdf_loaders/dotsocr.py | 4 +++- .../components/indexer/loaders/pdf_loaders/openai.py | 4 +++- openrag/core/config/__init__.py | 10 ++++++++-- 4 files changed, 16 insertions(+), 6 deletions(-) diff --git a/openrag/app_front.py b/openrag/app_front.py index 607f837d7..9674d2f6a 100644 --- a/openrag/app_front.py +++ b/openrag/app_front.py @@ -9,7 +9,7 @@ from chainlit.config import config as cl_config from chainlit.context import get_context from consts import PARTITION_PREFIX -from core.utils.logging import get_logger +from core.utils.logging import get_logger, mask_email from dotenv import load_dotenv from openai import AsyncOpenAI @@ -130,7 +130,7 @@ async def auth_callback(username: str, password: str): ) except httpx.HTTPStatusError: - logger.info("Authentication failed", username=username) + logger.info("Authentication failed", username=mask_email(username)) return None except Exception as e: logger.exception("Unexpected error during authentication", error=str(e)) diff --git a/openrag/components/indexer/loaders/pdf_loaders/dotsocr.py b/openrag/components/indexer/loaders/pdf_loaders/dotsocr.py index c72c3638c..7c1e7a6af 100644 --- a/openrag/components/indexer/loaders/pdf_loaders/dotsocr.py +++ b/openrag/components/indexer/loaders/pdf_loaders/dotsocr.py @@ -1,9 +1,11 @@ -from core.utils.logging import logger # assuming you have a shared logger instance +from core.utils.logging import get_logger from PIL import Image from tqdm.asyncio import tqdm from .openai import OpenAILoader +logger = get_logger() + class DotsOCRLoader(OpenAILoader): """PDF loader using DotsOCR""" diff --git a/openrag/components/indexer/loaders/pdf_loaders/openai.py b/openrag/components/indexer/loaders/pdf_loaders/openai.py index 3ad7f6782..2ad4461c8 100644 --- a/openrag/components/indexer/loaders/pdf_loaders/openai.py +++ b/openrag/components/indexer/loaders/pdf_loaders/openai.py @@ -7,13 +7,15 @@ from pathlib import Path import pypdfium2 as pdfium -from core.utils.logging import logger +from core.utils.logging import get_logger from langchain.schema import Document from langchain_openai import ChatOpenAI from PIL import Image from ..base import BaseLoader +logger = get_logger() + async def pdf_to_images(pdf_path: str, scale: float = 1.0) -> list[Image.Image]: pdf: pdfium.PdfDocument = await asyncio.to_thread(pdfium.PdfDocument, pdf_path) diff --git a/openrag/core/config/__init__.py b/openrag/core/config/__init__.py index 7f91fb315..014e9d6ec 100644 --- a/openrag/core/config/__init__.py +++ b/openrag/core/config/__init__.py @@ -6,7 +6,10 @@ get_settings() — cached singleton accessor """ +from collections.abc import Mapping from functools import lru_cache +from pathlib import Path +from typing import Any from .root import Settings @@ -19,7 +22,10 @@ def get_settings() -> Settings: return _load() -def load_config(config_path=None, overrides=None) -> Settings: +def load_config( + config_path: str | Path | None = None, + overrides: Mapping[str, Any] | None = None, +) -> Settings: """Return the cached Pydantic Settings singleton. The ``config_path`` parameter is kept for backward compatibility. @@ -27,7 +33,7 @@ def load_config(config_path=None, overrides=None) -> Settings: The ``overrides`` parameter bypasses the cache (useful for tests). """ - if overrides or config_path: + if overrides is not None or config_path is not None: from .loader import load_config as _load return _load(conf_dir=config_path, overrides=overrides) From c7be6380ed621047e8345ea95e455a6b01ff0d7d Mon Sep 17 00:00:00 2001 From: hedhoud <74668966+hedhoud@users.noreply.github.com> Date: Fri, 29 May 2026 14:58:49 +0200 Subject: [PATCH 04/11] refactor: move utility and chunker helpers to canonical layers --- openrag/api/dependencies/files.py | 26 +++ .../utils => api/dependencies}/test_files.py | 15 +- openrag/api/routers/admin/indexing.py | 3 +- openrag/api/routers/admin/tools.py | 2 +- openrag/api/routers/user/chat.py | 3 +- openrag/core/chunking/factory.py | 26 +++ openrag/core/utils/filename.py | 25 +++ openrag/core/utils/source_filtering.py | 153 ++++++++++++++++++ .../utils}/test_source_filtering.py | 5 +- openrag/core/utils/text.py | 30 ++++ .../inference/parsers/openai_audio.py | 2 +- openrag/services/inference/runtime.py | 47 ++++++ .../orchestrators/indexing_service.py | 2 +- .../services/orchestrators/query_service.py | 42 +++-- openrag/services/workers/indexer_pool.py | 11 +- openrag/services/workers/test_indexer_pool.py | 27 +--- 16 files changed, 362 insertions(+), 57 deletions(-) rename openrag/{components/indexer/utils => api/dependencies}/test_files.py (88%) create mode 100644 openrag/core/chunking/factory.py create mode 100644 openrag/core/utils/source_filtering.py rename openrag/{components => core/utils}/test_source_filtering.py (99%) create mode 100644 openrag/services/inference/runtime.py diff --git a/openrag/api/dependencies/files.py b/openrag/api/dependencies/files.py index a6e3f2c03..5f262a13b 100644 --- a/openrag/api/dependencies/files.py +++ b/openrag/api/dependencies/files.py @@ -1,6 +1,10 @@ +from pathlib import Path from typing import Any +import aiofiles +import consts from core.indexing import validators as core_validators +from core.utils.filename import make_unique_filename from di.providers import get_config from fastapi import Depends, Form, UploadFile @@ -29,3 +33,25 @@ async def validate_file_format( mimetype=metadata.get("mimetype"), ) return file + + +async def save_file_to_disk( + file: UploadFile, + dest_dir: Path, + chunk_size: int = consts.FILE_READ_CHUNK_SIZE, + with_random_prefix: bool = False, +) -> Path: + """Save an uploaded file to disk in chunks and return the saved path.""" + dest_dir.mkdir(parents=True, exist_ok=True) + + filename = make_unique_filename(file.filename) if with_random_prefix else file.filename + file_path = dest_dir / filename + + async with aiofiles.open(file_path, "wb") as buffer: + while True: + chunk = await file.read(chunk_size) + if not chunk: + break + await buffer.write(chunk) + + return file_path diff --git a/openrag/components/indexer/utils/test_files.py b/openrag/api/dependencies/test_files.py similarity index 88% rename from openrag/components/indexer/utils/test_files.py rename to openrag/api/dependencies/test_files.py index ff46da5fd..ab5b9e40b 100644 --- a/openrag/components/indexer/utils/test_files.py +++ b/openrag/api/dependencies/test_files.py @@ -2,9 +2,10 @@ from pathlib import Path import pytest -from fastapi import HTTPException, UploadFile - -from .files import extract_temporal_fields, sanitize_filename, save_file_to_disk +from api.dependencies.files import save_file_to_disk +from core.utils.exceptions import ValidationError +from core.utils.filename import extract_temporal_fields, sanitize_filename +from fastapi import UploadFile @pytest.mark.asyncio @@ -36,7 +37,7 @@ def fake_make_unique_filename(filename: str) -> str: return "PREFIX_1234_test.txt" monkeypatch.setattr( - "openrag.components.indexer.utils.files.make_unique_filename", + "api.dependencies.files.make_unique_filename", fake_make_unique_filename, ) @@ -105,8 +106,8 @@ def test_extract_temporal_fields_with_timezone(): def test_extract_temporal_fields_invalid_datetime_raises_400(): - with pytest.raises(HTTPException) as exc_info: + with pytest.raises(ValidationError) as exc_info: extract_temporal_fields({"created_at": "not-a-date"}, ["created_at"]) assert exc_info.value.status_code == 400 - assert "not-a-date" in exc_info.value.detail - assert "created_at" in exc_info.value.detail + assert "not-a-date" in str(exc_info.value) + assert "created_at" in str(exc_info.value) diff --git a/openrag/api/routers/admin/indexing.py b/openrag/api/routers/admin/indexing.py index e43a354af..64b640450 100644 --- a/openrag/api/routers/admin/indexing.py +++ b/openrag/api/routers/admin/indexing.py @@ -23,12 +23,13 @@ require_task_owner, ) from api.dependencies.files import ( + save_file_to_disk, validate_file_format, validate_file_id, validate_metadata, ) from api.routers.admin.task_logs import collect_task_logs -from components.indexer.utils.files import sanitize_filename, save_file_to_disk +from core.utils.filename import sanitize_filename from core.utils.log_tail import app_log_file from core.utils.logging import get_logger from di.providers import get_auth_service, get_config, get_indexing_service, get_partition_service diff --git a/openrag/api/routers/admin/tools.py b/openrag/api/routers/admin/tools.py index 34598a18d..ede394a6b 100644 --- a/openrag/api/routers/admin/tools.py +++ b/openrag/api/routers/admin/tools.py @@ -13,11 +13,11 @@ from pathlib import Path from api.dependencies.files import ( + save_file_to_disk, validate_file_format, validate_metadata, ) from api.schemas.admin.tools import ToolInfo -from components.indexer.utils.files import save_file_to_disk from core.utils.logging import get_logger from di.providers import get_config, get_conversion_service from fastapi import APIRouter, Depends, Form, HTTPException, UploadFile, status diff --git a/openrag/api/routers/user/chat.py b/openrag/api/routers/user/chat.py index 3e57ba8f7..03badfc18 100644 --- a/openrag/api/routers/user/chat.py +++ b/openrag/api/routers/user/chat.py @@ -30,11 +30,10 @@ ) from api.routers.user.source_links import build_document_source_link from api.schemas.user.chat import OpenAIChatCompletionRequest, OpenAICompletionRequest -from components.indexer.utils.text_sanitizer import sanitize_text -from components.utils import get_num_tokens from config import load_config from core.utils.exceptions import OpenRAGError from core.utils.logging import get_logger +from core.utils.text import get_num_tokens, sanitize_text from di.providers import get_config, get_partition_service, get_query_service from fastapi import APIRouter, Body, Depends, HTTPException, Request, status from fastapi.responses import JSONResponse, StreamingResponse diff --git a/openrag/core/chunking/factory.py b/openrag/core/chunking/factory.py new file mode 100644 index 000000000..33a58d07a --- /dev/null +++ b/openrag/core/chunking/factory.py @@ -0,0 +1,26 @@ +"""Factory for configured chunking strategies.""" + +from typing import Any + +import core.chunking.recursive # noqa: F401 +from core.chunking.chunking_strategy import ChunkingStrategy +from core.chunking.registry import chunking_registry +from core.utils.exceptions import RegistryError +from core.utils.text import get_num_tokens + + +def create_chunker(config: Any) -> ChunkingStrategy: + """Create the configured chunking strategy.""" + chunker_params = config.chunker.model_dump() + name = chunker_params.pop("name") + + try: + return chunking_registry.create( + name, + length_function=get_num_tokens(), + **chunker_params, + ) + except RegistryError as exc: + raise ValueError( + f"Chunker '{name}' is not recognized. Available chunkers: {chunking_registry.list_registered()}" + ) from exc diff --git a/openrag/core/utils/filename.py b/openrag/core/utils/filename.py index 7fe852970..2eed41e83 100644 --- a/openrag/core/utils/filename.py +++ b/openrag/core/utils/filename.py @@ -8,8 +8,11 @@ import re import secrets import time +from datetime import UTC, datetime from pathlib import Path +from core.utils.exceptions import ValidationError + def sanitize_filename(filename: str) -> str: """Sanitize a filename by removing special characters. @@ -47,3 +50,25 @@ def make_unique_filename(filename: str) -> str: ts = int(time.time() * 1000) rand = secrets.token_hex(2) return f"{ts}_{rand}_{filename}" + + +def extract_temporal_fields(metadata: dict, temporal_fields: list) -> dict: + """Extract and validate ISO-8601 temporal metadata fields.""" + result = {} + for field in temporal_fields: + if field not in metadata or metadata[field] is None: + continue + + datetime_str = metadata[field] + try: + parsed = datetime.fromisoformat(datetime_str) + if parsed.tzinfo is None: + parsed = parsed.replace(tzinfo=UTC) + result[field] = parsed.isoformat() + except Exception: + raise ValidationError( + f"Invalid ISO 8601 datetime field ({datetime_str}) for field '{field}'.", + status_code=400, + ) + + return result diff --git a/openrag/core/utils/source_filtering.py b/openrag/core/utils/source_filtering.py new file mode 100644 index 000000000..ff3d674fe --- /dev/null +++ b/openrag/core/utils/source_filtering.py @@ -0,0 +1,153 @@ +"""Source citation extraction and streaming filtering helpers.""" + +from __future__ import annotations + +import asyncio +import copy +import json +import re + +from core.utils.logging import get_logger + +logger = get_logger() + +_SOURCES_NONE_RE = re.compile( + r"\n?[ \t]*\[?Sources?\]?\s*:\s*\[?\s*none\s*\]?[.\s]*?(?=\n|$)", + re.IGNORECASE, +) +_SOURCES_NUMS_RE = re.compile(r"\n?[ \t]*\[?Sources?\]?\s*:\s*\[?([\d,\s]+)\]?[.\s]*?(?=\n|$)") + + +def _strip_sources_tags(text: str) -> tuple[str, set[int], bool]: + """Strip line-terminal source tags and return citations found.""" + cited: set[int] = set() + for match in _SOURCES_NUMS_RE.finditer(text): + cited.update(int(n.strip()) for n in match.group(1).split(",") if n.strip().isdigit()) + saw_none = bool(_SOURCES_NONE_RE.search(text)) + cleaned = _SOURCES_NUMS_RE.sub("", text) + cleaned = _SOURCES_NONE_RE.sub("", cleaned) + return cleaned, cited, saw_none + + +def extract_and_strip_sources_block(text: str) -> tuple[str, set[int] | None]: + """Strip line-terminal source tags and return merged citations.""" + cleaned, citations, saw_none = _strip_sources_tags(text) + + if not citations and not saw_none: + tail = text[-150:] if len(text) > 150 else text + logger.debug("No [Sources: ...] tag found in LLM response", tail=repr(tail)) + return text, None + + cleaned = cleaned.rstrip() + if citations: + logger.debug("Extracted source citations from LLM response", citations=sorted(citations)) + return cleaned, citations + + logger.debug("LLM explicitly reported no sources used") + return cleaned, set() + + +def filter_sources_by_citations(sources: list, citations: set[int] | None) -> list: + """Keep only sources whose 1-based index was cited.""" + if citations is None: + return sources + if not citations: + return [] + filtered = [source for i, source in enumerate(sources, start=1) if i in citations] + return filtered if filtered else sources + + +def _min_sources_tag_buffer_size(n_sources: int) -> int: + """Pessimistic upper bound on the length of a ``[Sources: ...]`` tag.""" + if n_sources <= 0: + return 100 + digits_total = sum(len(str(i)) for i in range(1, n_sources + 1)) + separators = max(0, n_sources - 1) * 2 + wrapping = len("\n[Sources: ") + len("]") + 8 + return digits_total + separators + wrapping + + +_MIN_STREAM_LOOKAHEAD = 80 + + +async def stream_with_source_filtering( + llm_stream, + sources: list, + model_name: str, + buffer_size: int | None = None, +): + """Process an LLM SSE stream, stripping line-terminal source tags.""" + if buffer_size is None: + buffer_size = max(_MIN_STREAM_LOOKAHEAD, _min_sources_tag_buffer_size(len(sources))) + pending = "" + emitted_len = 0 + chunk_template = None + last_finish_reason = None + + async for line in llm_stream: + if not line.startswith("data:"): + continue + + if line.strip() == "data: [DONE]": + final_clean, citations = extract_and_strip_sources_block(pending) + final_clean = final_clean.rstrip() + + filtered = filter_sources_by_citations(sources, citations) + filtered_json = json.dumps({"sources": filtered}) + + if chunk_template and len(final_clean) > emitted_len: + tail_chunk = copy.deepcopy(chunk_template) + tail_chunk["choices"][0]["delta"] = {"content": final_clean[emitted_len:]} + tail_chunk["extra"] = filtered_json + yield f"data: {json.dumps(tail_chunk)}\n\n" + + if chunk_template: + await asyncio.sleep(0.05) + finish_chunk = copy.deepcopy(chunk_template) + finish_chunk["choices"][0]["delta"] = {} + finish_chunk["choices"][0]["finish_reason"] = last_finish_reason or "stop" + finish_chunk["extra"] = filtered_json + yield f"data: {json.dumps(finish_chunk)}\n\n" + + yield "data: [DONE]\n\n" + continue + + data = json.loads(line[len("data: ") :]) + data["model"] = model_name + + choice = data.get("choices", [{}])[0] + delta = choice.get("delta", {}) + content = delta.get("content", "") or "" + finish_reason = choice.get("finish_reason") + + if finish_reason: + last_finish_reason = finish_reason + chunk_template = data + elif content: + chunk_template = data + pending += content + + if len(pending) <= buffer_size: + continue + + cleaned, _, _ = _strip_sources_tags(pending) + safe_end = max(0, len(cleaned) - buffer_size) + if safe_end > emitted_len: + out = { + **data, + "choices": [ + { + **choice, + "delta": { + **choice.get("delta", {}), + "content": cleaned[emitted_len:safe_end], + }, + } + ], + "extra": "{}", + } + yield f"data: {json.dumps(out)}\n\n" + emitted_len = safe_end + else: + data["extra"] = "{}" + yield f"data: {json.dumps(data)}\n\n" diff --git a/openrag/components/test_source_filtering.py b/openrag/core/utils/test_source_filtering.py similarity index 99% rename from openrag/components/test_source_filtering.py rename to openrag/core/utils/test_source_filtering.py index 11bd2f80f..9140886ad 100644 --- a/openrag/components/test_source_filtering.py +++ b/openrag/core/utils/test_source_filtering.py @@ -3,7 +3,8 @@ import json import pytest -from components.utils import ( +from core.utils.source_filtering import ( + _min_sources_tag_buffer_size, extract_and_strip_sources_block, filter_sources_by_citations, stream_with_source_filtering, @@ -321,8 +322,6 @@ async def test_mid_response_tag_stripped_plus_trailing_tag(self): def test_min_sources_tag_buffer_size_fits_many_sources(): - from components.utils import _min_sources_tag_buffer_size - for n in (1, 10, 60, 100): tag = "\n[Sources: " + ", ".join(str(i) for i in range(1, n + 1)) + "]" assert _min_sources_tag_buffer_size(n) >= len(tag), n diff --git a/openrag/core/utils/text.py b/openrag/core/utils/text.py index b992191f6..656d39f68 100644 --- a/openrag/core/utils/text.py +++ b/openrag/core/utils/text.py @@ -9,8 +9,38 @@ import re import unicodedata +from core.config import load_config +from core.utils.logging import get_logger + DEFAULT_FALLBACK_ENCODING = "utf-8" +logger = get_logger() + + +_cached_length_function = None + + +def get_num_tokens(): + """Return the configured token counter, with a local tiktoken fallback.""" + global _cached_length_function + if _cached_length_function is None: + try: + from langchain_openai import ChatOpenAI + + config = load_config() + llm = ChatOpenAI(**config.llm.model_dump()) + _cached_length_function = llm.get_num_tokens + except Exception as exc: + import tiktoken + + logger.warning( + "ChatOpenAI unavailable for token counting, falling back to tiktoken cl100k_base", + error=str(exc), + ) + encoding = tiktoken.get_encoding("cl100k_base") + _cached_length_function = lambda text: len(encoding.encode(text)) # noqa: E731 + return _cached_length_function + def decode_bytes(raw: bytes, encoding: str | None = None) -> str: """Decode ``raw`` to ``str`` with a UTF-8-first detection strategy. diff --git a/openrag/services/inference/parsers/openai_audio.py b/openrag/services/inference/parsers/openai_audio.py index 54058d0ad..af34e1702 100644 --- a/openrag/services/inference/parsers/openai_audio.py +++ b/openrag/services/inference/parsers/openai_audio.py @@ -17,7 +17,7 @@ Adapted from the legacy ``components/indexer/loaders/audio/openai.py`` ``AudioTranscriber``; -the new version drops the in-memory ``components.utils`` semaphore (now +the new version drops the old in-memory semaphore helper (now per-instance via ``concurrency_limit``) and the embedded WhisperActor ref-getter (now an injected callable). """ diff --git a/openrag/services/inference/runtime.py b/openrag/services/inference/runtime.py new file mode 100644 index 000000000..166ed3b92 --- /dev/null +++ b/openrag/services/inference/runtime.py @@ -0,0 +1,47 @@ +"""Runtime inference helpers that depend on infrastructure services.""" + +from core.config import load_config +from fast_langdetect import LangDetectConfig, LangDetector +from services.inference.distributed_semaphore import DistributedSemaphore + +_LANG_DETECT_CACHE_DIR = "/app/model_weights/" +_lang_detector = LangDetector( + config=LangDetectConfig( + max_input_length=1024, + model="auto", + cache_dir=_LANG_DETECT_CACHE_DIR, + ) +) + + +def detect_language(text: str): + """Detect the primary language of ``text``.""" + outputs = _lang_detector.detect(text, k=1) + return outputs[0].get("lang") + + +def get_llm_semaphore() -> DistributedSemaphore: + """Return the distributed semaphore for LLM calls.""" + config = load_config() + return DistributedSemaphore( + name="llmSemaphore", + max_concurrent_ops=config.semaphore.llm_semaphore, + ) + + +def get_vlm_semaphore() -> DistributedSemaphore: + """Return the distributed semaphore for VLM calls.""" + config = load_config() + return DistributedSemaphore( + name="vlmSemaphore", + max_concurrent_ops=config.semaphore.vlm_semaphore, + ) + + +def get_audio_semaphore() -> DistributedSemaphore: + """Return the distributed semaphore for audio transcription calls.""" + config = load_config() + return DistributedSemaphore( + name="audioSemaphore", + max_concurrent_ops=config.loader.transcriber.max_concurrent_chunks, + ) diff --git a/openrag/services/orchestrators/indexing_service.py b/openrag/services/orchestrators/indexing_service.py index 9095efab1..3850bd9c8 100644 --- a/openrag/services/orchestrators/indexing_service.py +++ b/openrag/services/orchestrators/indexing_service.py @@ -16,7 +16,7 @@ from pathlib import Path from typing import TYPE_CHECKING -from components.indexer.utils.files import extract_temporal_fields +from core.utils.filename import extract_temporal_fields from core.utils.logging import get_logger if TYPE_CHECKING: diff --git a/openrag/services/orchestrators/query_service.py b/openrag/services/orchestrators/query_service.py index 20f81e4c0..6a8d8d643 100644 --- a/openrag/services/orchestrators/query_service.py +++ b/openrag/services/orchestrators/query_service.py @@ -16,7 +16,7 @@ (retry → raw user query; relevancy=False on parse failure). * **Streaming + citations live here; the router is pure transport.** ``chat_stream`` drives the proven - ``components.utils.stream_with_source_filtering`` (100-char buffer that + ``core.utils.source_filtering.stream_with_source_filtering`` (100-char buffer that strips the ``[Sources: N]`` tag before it reaches the client); ``chat`` / ``complete`` return the finalized OpenAI dict with the citation-filtered ``extra`` sources. The router only maps the @@ -24,7 +24,7 @@ callable — keeps ``request.url_for`` in transport), and wraps ``StreamingResponse`` / ``JSONResponse``. -Imports from ``components.*`` (pure helpers / prompts / websearch) are +Imports from ``components.*`` (prompt shims) are allowed during the Phase-8 shim (legacy layer, unchecked by the guard; no LangChain symbol is imported into this file → 8H clean). ``Chunk`` is converted to LangChain ``Document`` via ``Chunk.to_langchain()`` at the @@ -47,18 +47,20 @@ SPOKEN_STYLE_ANSWER_PROMPT, SYS_PROMPT_TMPLT, ) -from components.utils import ( +from core.models.query import Query, SearchQueries +from core.prompts import ( SOURCE_SEPARATOR, - detect_language, - extract_and_strip_sources_block, - filter_sources_by_citations, format_context, format_web_context, - get_llm_semaphore, - stream_with_source_filtering, ) -from core.models.query import Query, SearchQueries from core.utils.logging import get_logger +from core.utils.source_filtering import ( + extract_and_strip_sources_block, + filter_sources_by_citations, + stream_with_source_filtering, +) +from core.utils.text import get_num_tokens +from services.inference.runtime import detect_language, get_llm_semaphore if TYPE_CHECKING: from core.config.root import Settings @@ -257,15 +259,25 @@ async def _prepare_chat(self, partition: list[str] | None, payload: dict): web_formatted, web_tokens = "", 0 if web_results: web_formatted, _, web_tokens = format_web_context( - web_results, start_index=1, max_tokens=self._web.max_tokens + web_results, + length_function=get_num_tokens(), + start_index=1, + max_tokens=self._web.max_tokens, ) - context, included = format_context(docs, max_context_tokens=self._max_context_tokens - web_tokens) + context, included = format_context( + [doc.page_content for doc in docs], + max_context_tokens=self._max_context_tokens - web_tokens, + length_function=get_num_tokens(), + ) docs = [docs[i] for i in included] if web_results: if docs: web_formatted, _, _ = format_web_context( - web_results, start_index=len(docs) + 1, max_tokens=self._web.max_tokens + web_results, + length_function=get_num_tokens(), + start_index=len(docs) + 1, + max_tokens=self._web.max_tokens, ) else: context = "" @@ -298,7 +310,11 @@ async def _prepare_completions(self, partition: list[str], payload: dict): queries = await self.generate_query([{"role": "user", "content": prompt}]) chunks = await self._retrieval.retrieve_multi(partitions=partition, search_queries=queries) docs = [c.to_langchain() for c in chunks] - context, included = format_context(docs, max_context_tokens=self._max_context_tokens) + context, included = format_context( + [doc.page_content for doc in docs], + max_context_tokens=self._max_context_tokens, + length_function=get_num_tokens(), + ) docs = [docs[i] for i in included] if docs: payload["prompt"] = ( diff --git a/openrag/services/workers/indexer_pool.py b/openrag/services/workers/indexer_pool.py index 54d74e72b..3b7cafca9 100644 --- a/openrag/services/workers/indexer_pool.py +++ b/openrag/services/workers/indexer_pool.py @@ -99,15 +99,12 @@ def build_indexer_pool(namespace: str = "openrag") -> Any: def _build_chunker(cfg: Any) -> Any: - from components.indexer.chunker.chunker import ChunkerFactory + from core.chunking.factory import create_chunker - legacy_chunker = ChunkerFactory.create_chunker(cfg) - if hasattr(legacy_chunker, "chunk"): - return legacy_chunker - core_chunker = getattr(legacy_chunker, "_core_splitter", None) - if core_chunker is None or not hasattr(core_chunker, "chunk"): + chunker = create_chunker(cfg) + if not hasattr(chunker, "chunk"): raise TypeError("Configured chunker does not expose a chunk(document, partition) method") - return core_chunker + return chunker __all__ = ["IndexerPool", "build_indexer_pool"] diff --git a/openrag/services/workers/test_indexer_pool.py b/openrag/services/workers/test_indexer_pool.py index 127d95acc..31129090d 100644 --- a/openrag/services/workers/test_indexer_pool.py +++ b/openrag/services/workers/test_indexer_pool.py @@ -8,40 +8,25 @@ def chunk(self, document, partition: str = "default"): return [] -class _LegacyChunker: - def __init__(self) -> None: - self._core_splitter = _NativeChunker() - - -class _BrokenLegacyChunker: +class _BrokenChunker: pass def test_build_chunker_returns_native_chunker(monkeypatch: pytest.MonkeyPatch) -> None: - from components.indexer.chunker.chunker import ChunkerFactory + import core.chunking.factory as factory from services.workers.indexer_pool import _build_chunker native = _NativeChunker() - monkeypatch.setattr(ChunkerFactory, "create_chunker", staticmethod(lambda _cfg: native)) + monkeypatch.setattr(factory, "create_chunker", lambda _cfg: native) assert _build_chunker(object()) is native -def test_build_chunker_unwraps_legacy_core_splitter(monkeypatch: pytest.MonkeyPatch) -> None: - from components.indexer.chunker.chunker import ChunkerFactory - from services.workers.indexer_pool import _build_chunker - - legacy = _LegacyChunker() - monkeypatch.setattr(ChunkerFactory, "create_chunker", staticmethod(lambda _cfg: legacy)) - - assert _build_chunker(object()) is legacy._core_splitter - - -def test_build_chunker_rejects_invalid_legacy_chunker(monkeypatch: pytest.MonkeyPatch) -> None: - from components.indexer.chunker.chunker import ChunkerFactory +def test_build_chunker_rejects_invalid_chunker(monkeypatch: pytest.MonkeyPatch) -> None: + import core.chunking.factory as factory from services.workers.indexer_pool import _build_chunker - monkeypatch.setattr(ChunkerFactory, "create_chunker", staticmethod(lambda _cfg: _BrokenLegacyChunker())) + monkeypatch.setattr(factory, "create_chunker", lambda _cfg: _BrokenChunker()) with pytest.raises(TypeError, match="chunk"): _build_chunker(object()) From 3f5e462b58136a82ad384368edae50523ef2eb00 Mon Sep 17 00:00:00 2001 From: cbizeul Date: Fri, 29 May 2026 13:05:43 +0000 Subject: [PATCH 05/11] refactor: switch auth imports from components.auth to services.auth Phase 12B. The canonical auth adapters (OIDC client, session tokens, state cookie, deps) live in services/auth/; components/auth/ is now only re-export shims. Repoint the remaining consumers off the shims: - api/middleware/auth.py components.auth.refresh -> services.auth.refresh - api/routers/auth/oidc.py components.auth -> services.auth.state_cookie - di/container.py components.auth -> services.auth - services/orchestrators/auth_service.py components.auth -> services.auth - services/orchestrators/test_auth_service.py components.auth -> services.auth Also refresh stale module references in docstrings that pointed at the old components.auth paths (auth_service header, oidc_session_repo notes). No remaining `from components.auth` imports outside components/ itself. --- openrag/api/middleware/auth.py | 2 +- openrag/api/routers/auth/oidc.py | 2 +- openrag/di/container.py | 2 +- openrag/services/orchestrators/auth_service.py | 12 +++++------- openrag/services/orchestrators/test_auth_service.py | 4 ++-- openrag/services/persistence/oidc_session_repo.py | 4 ++-- 6 files changed, 12 insertions(+), 14 deletions(-) diff --git a/openrag/api/middleware/auth.py b/openrag/api/middleware/auth.py index 6a87f06b6..fdb9adff4 100644 --- a/openrag/api/middleware/auth.py +++ b/openrag/api/middleware/auth.py @@ -34,11 +34,11 @@ from typing import Any from urllib.parse import quote -from components.auth.refresh import refresh_session_if_needed from core.config.auth import AuthBypassConfig from core.utils.logging import get_logger from fastapi import Request from fastapi.responses import JSONResponse, RedirectResponse +from services.auth.refresh import refresh_session_if_needed from starlette.middleware.base import BaseHTTPMiddleware logger = get_logger() diff --git a/openrag/api/routers/auth/oidc.py b/openrag/api/routers/auth/oidc.py index 2df4ff4d1..ec5e91dc6 100644 --- a/openrag/api/routers/auth/oidc.py +++ b/openrag/api/routers/auth/oidc.py @@ -23,12 +23,12 @@ import os -from components.auth import StateCookieSerializer from core.utils.exceptions import OpenRAGError from core.utils.logging import get_logger from di.providers import get_auth_service from fastapi import APIRouter, Depends, Form, HTTPException, Request, Response, status from fastapi.responses import JSONResponse, RedirectResponse +from services.auth.state_cookie import StateCookieSerializer logger = get_logger() router = APIRouter() diff --git a/openrag/di/container.py b/openrag/di/container.py index 66a4a46a0..fbba92550 100644 --- a/openrag/di/container.py +++ b/openrag/di/container.py @@ -276,7 +276,7 @@ def auth_service(self) -> AuthService: cfg = self._oidc_config client = None if cfg.enabled: - from components.auth import get_oidc_client + from services.auth import get_oidc_client client = get_oidc_client() self._auth_service = AuthService( diff --git a/openrag/services/orchestrators/auth_service.py b/openrag/services/orchestrators/auth_service.py index f14a8d4a6..8922c4b1a 100644 --- a/openrag/services/orchestrators/auth_service.py +++ b/openrag/services/orchestrators/auth_service.py @@ -12,9 +12,7 @@ The cryptographic / cookie primitives (``OIDCClient``, the state-cookie serializer, Fernet token (de)encryption, opaque session-token issuance) -still come from ``components.auth`` during the Phase-8 shim period — -those are infrastructure adapters scheduled to move under -``services/auth`` in Phase 9. +come from ``services.auth`` — the infrastructure adapters for OIDC. """ from __future__ import annotations @@ -25,7 +23,10 @@ from typing import TYPE_CHECKING, Any from urllib.parse import urlencode, urlparse -from components.auth import ( +from core.models.user import OIDCSession, User +from core.utils.exceptions import AuthError, OpenRAGError +from core.utils.logging import get_logger, mask_email +from services.auth import ( OIDCClient, StateCookiePayload, StateCookieSerializer, @@ -34,9 +35,6 @@ hash_session_token, issue_session_token, ) -from core.models.user import OIDCSession, User -from core.utils.exceptions import AuthError, OpenRAGError -from core.utils.logging import get_logger, mask_email if TYPE_CHECKING: from core.config.auth import OIDCConfig diff --git a/openrag/services/orchestrators/test_auth_service.py b/openrag/services/orchestrators/test_auth_service.py index 527f00b5f..e7fda1a29 100644 --- a/openrag/services/orchestrators/test_auth_service.py +++ b/openrag/services/orchestrators/test_auth_service.py @@ -9,11 +9,11 @@ from __future__ import annotations import pytest -from components.auth import StateCookieSerializer, hash_session_token -from components.auth.oidc_client import LogoutTokenClaims, TokenBundle from core.config.auth import OIDCConfig from core.models.user import User from cryptography.fernet import Fernet +from services.auth import StateCookieSerializer, hash_session_token +from services.auth.oidc_client import LogoutTokenClaims, TokenBundle from services.orchestrators.auth_service import AuthService, OIDCFlowError KEY = Fernet.generate_key().decode() diff --git a/openrag/services/persistence/oidc_session_repo.py b/openrag/services/persistence/oidc_session_repo.py index 482d97728..4de5e335d 100644 --- a/openrag/services/persistence/oidc_session_repo.py +++ b/openrag/services/persistence/oidc_session_repo.py @@ -10,7 +10,7 @@ Tokens (``id_token``, ``access_token``, ``refresh_token``) are stored **Fernet-encrypted** as ``BYTEA``. Encryption and decryption are the -caller's responsibility (see ``components.auth.crypto``) — the repo +caller's responsibility (see ``services.auth.session_tokens``) — the repo treats the bytes as opaque. The plain session cookie value is hashed (SHA-256) at the caller before being passed in. @@ -302,7 +302,7 @@ async def update_oidc_session_tokens( ``SELECT ... FOR UPDATE`` serialises concurrent refresh callers on the same session row — Postgres only. The wider stampede guard in - :mod:`components.auth.refresh` short-circuits before this is even + :mod:`services.auth.refresh` short-circuits before this is even called in the common case. """ async with self.pool.acquire() as conn: From 6eda2b6cc0ff0de42877c7bcf32e32e4c33fa8ae Mon Sep 17 00:00:00 2001 From: cbizeul Date: Fri, 29 May 2026 13:27:57 +0000 Subject: [PATCH 06/11] refactor: load retrieval/query prompts via core loader, drop components.prompts MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Phase 12C. The orchestrators imported HYDE/MULTI_QUERY/SYS/CONTEXTUALIZER/ SPOKEN_STYLE prompts from the components.prompts shim. Those are not Python string constants — they are disk templates the shim eager-loads at import via load_config(). Distributing them as eager module-level constants into core/prompts builders would add import-time disk I/O and a config dependency to the whole core.prompts package (its __init__ imports every builder), so instead each service loads its templates from the injected Settings using the existing pure core.prompts.load_template_by_key(prompts_dir, mapping, key) — exactly the call shape template_loader documents for its callers. - retrieval_service.py: multiQuery / hyde branches load "multi_query" / "hyde" from `config` (already in __init__) instead of importing the shim inline. - query_service.py: load "query_contextualizer" / "spoken_style_answer" / "sys_prompt" once in __init__ from the injected config; store on the instance. Per-instance config binding replaces the old global eager load. - test_query_service.py: extend the fake config with the real paths/prompts so QueryService can resolve templates. No remaining `components.prompts` imports outside components/ itself. --- openrag/services/orchestrators/query_service.py | 15 ++++++++------- .../services/orchestrators/retrieval_service.py | 9 +++------ .../services/orchestrators/test_query_service.py | 7 +++++++ 3 files changed, 18 insertions(+), 13 deletions(-) diff --git a/openrag/services/orchestrators/query_service.py b/openrag/services/orchestrators/query_service.py index 6a8d8d643..da4f92a77 100644 --- a/openrag/services/orchestrators/query_service.py +++ b/openrag/services/orchestrators/query_service.py @@ -42,16 +42,12 @@ from enum import Enum from typing import TYPE_CHECKING, Any -from components.prompts import ( - QUERY_CONTEXTUALIZER_PROMPT, - SPOKEN_STYLE_ANSWER_PROMPT, - SYS_PROMPT_TMPLT, -) from core.models.query import Query, SearchQueries from core.prompts import ( SOURCE_SEPARATOR, format_context, format_web_context, + load_template_by_key, ) from core.utils.logging import get_logger from core.utils.source_filtering import ( @@ -127,6 +123,11 @@ def __init__( self._mr_expansion = mr.expansion_batch_size self._mr_max = mr.max_total_documents + prompts_dir, mapping = config.paths.prompts_dir, config.prompts + self._query_contextualizer_prompt = load_template_by_key(prompts_dir, mapping, "query_contextualizer") + self._spoken_style_answer_prompt = load_template_by_key(prompts_dir, mapping, "spoken_style_answer") + self._sys_prompt_tmplt = load_template_by_key(prompts_dir, mapping, "sys_prompt") + # ------------------------------------------------------------------ # Query generation (was RagPipeline.generate_query — no LangChain) # ------------------------------------------------------------------ @@ -137,7 +138,7 @@ async def generate_query(self, messages: list[dict]) -> SearchQueries: return SearchQueries(query_list=[Query(query=last_user)]) chat_history = "".join(f"{m['role']}: {m['content']}\n" for m in messages) - prompt = QUERY_CONTEXTUALIZER_PROMPT.format( + prompt = self._query_contextualizer_prompt.format( query_language=detect_language(last_user), current_date=datetime.now().strftime("%A, %B %d, %Y, %H:%M:%S"), ) @@ -284,7 +285,7 @@ async def _prepare_chat(self, partition: list[str] | None, payload: dict): context = f"{context}{SOURCE_SEPARATOR}{web_formatted}" if context else web_formatted new_messages = copy.deepcopy(messages) - tmpl = SPOKEN_STYLE_ANSWER_PROMPT if spoken_style else SYS_PROMPT_TMPLT + tmpl = self._spoken_style_answer_prompt if spoken_style else self._sys_prompt_tmplt new_messages.insert( 0, { diff --git a/openrag/services/orchestrators/retrieval_service.py b/openrag/services/orchestrators/retrieval_service.py index bf6487d12..5ec81db3e 100644 --- a/openrag/services/orchestrators/retrieval_service.py +++ b/openrag/services/orchestrators/retrieval_service.py @@ -27,6 +27,7 @@ import asyncio from typing import TYPE_CHECKING +from core.prompts import load_template_by_key from core.retrieval.pipeline import RetrieverPipeline from core.retrieval.retriever import ( HyDeRetriever, @@ -77,20 +78,16 @@ def __init__( } rtype = rcfg.type if rtype == "multiQuery": - from components.prompts import MULTI_QUERY_PROMPT - retriever = MultiQueryRetriever( llm=llm, - multi_query_template=MULTI_QUERY_PROMPT, + multi_query_template=load_template_by_key(config.paths.prompts_dir, config.prompts, "multi_query"), k_queries=rcfg.k_queries, **common, ) elif rtype == "hyde": - from components.prompts import HYDE_PROMPT - retriever = HyDeRetriever( llm=llm, - hyde_template=HYDE_PROMPT, + hyde_template=load_template_by_key(config.paths.prompts_dir, config.prompts, "hyde"), combine=rcfg.combine, **common, ) diff --git a/openrag/services/orchestrators/test_query_service.py b/openrag/services/orchestrators/test_query_service.py index 99f533636..bbae81540 100644 --- a/openrag/services/orchestrators/test_query_service.py +++ b/openrag/services/orchestrators/test_query_service.py @@ -15,9 +15,14 @@ import pytest import services.orchestrators.query_service as qs +from core.config import load_config from core.models.chunk import Chunk from services.orchestrators.query_service import QueryService +# Real prompt-template config (dir + key->filename mapping) so QueryService +# can load its system / contextualizer / spoken-style templates from disk. +_PROMPT_CFG = load_config() + @pytest.fixture(autouse=True) def _patch_infra(monkeypatch): @@ -88,6 +93,8 @@ def _config(mode="SimpleRag"): reranker=SimpleNamespace(top_k=5), chunker=SimpleNamespace(chunk_size=512), map_reduce=SimpleNamespace(initial_batch_size=2, expansion_batch_size=2, max_total_documents=4), + paths=_PROMPT_CFG.paths, + prompts=_PROMPT_CFG.prompts, ) From b750432479859fb3b7203e2a5e1348d7d8cf03d2 Mon Sep 17 00:00:00 2001 From: cbizeul Date: Fri, 29 May 2026 13:51:49 +0000 Subject: [PATCH 07/11] refactor: relocate legacy loader subsystem to services, repoint workers MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Phase 12F. Two parts. 1. Bootstrap re-points (safe — old modules are re-export shims): - components.indexer.loaders.serializer.DocSerializer -> services.workers.parsers.doc_serializer - components.indexer.loaders.pdf_loaders.docling2.DoclingPool -> services.workers.parsers.docling_workers - components.indexer.loaders.pdf_loaders.marker.MarkerPool -> services.workers.parsers.marker_workers 2. Loader-registry relocation. The plan assumed the loaders were already migrated to core/indexing/parsers/ and the old ones were dead. They are not: core/indexing/parsers/registry.py is an empty Registry, and the new doc_serializer / doc_serializer_bridge still drive the legacy loaders via get_loader_classes (which walks the loaders package and discovers ~20 BaseLoader adapters). So the legacy loader subsystem is the live runtime path. To make components/ deletable (12H) it is moved wholesale: components/indexer/loaders/ -> services/workers/parsers/legacy_loaders/ - get_loader_classes root_pkg updated to the new package path. - The two get_loader_classes consumers (doc_serializer, doc_serializer_bridge) and internal absolute imports / tests repointed. - The orphaned serializer.py re-export shim (no remaining consumers after the bootstrap re-point) is deleted rather than carried into services/. - Internal `from components.utils` / `from components.prompts` imports inside the loaders are left as-is; they remain valid and are 12C/12E/12H's concern. Runtime discovery verified: get_loader_classes resolves all 19 extension mappings from the new location. Also fix a latent test-isolation bug this relocation surfaced: services/inference/parsers/test_openai_audio.py installed a fake non-package `pydub` into sys.modules at import time and never restored it. It was masked only by collection order (the real-pydub loader test used to sort before it under components/). It now tries the real pydub first and only falls back to the stub when the import genuinely fails — its stated Python-3.13 purpose. No remaining `components.indexer.loaders` references anywhere. --- openrag/components/indexer/loaders/serializer.py | 11 ----------- .../services/inference/parsers/test_openai_audio.py | 9 ++++++--- openrag/services/workers/bootstrap.py | 6 +++--- openrag/services/workers/parsers/doc_serializer.py | 2 +- .../services/workers/parsers/doc_serializer_bridge.py | 2 +- .../parsers/legacy_loaders}/CustomDocLoader.py | 0 .../parsers/legacy_loaders}/CustomHTMLLoader.py | 0 .../workers/parsers/legacy_loaders}/__init__.py | 2 +- .../workers/parsers/legacy_loaders}/audio/__init__.py | 0 .../parsers/legacy_loaders}/audio/local_whisper.py | 2 +- .../workers/parsers/legacy_loaders}/audio/openai.py | 0 .../parsers/legacy_loaders}/audio/test_openai.py | 0 .../workers/parsers/legacy_loaders}/base.py | 0 .../workers/parsers/legacy_loaders}/doc.py | 0 .../workers/parsers/legacy_loaders}/docx.py | 0 .../workers/parsers/legacy_loaders}/eml_loader.py | 0 .../workers/parsers/legacy_loaders}/image.py | 0 .../parsers/legacy_loaders}/pdf_loaders/__init__.py | 0 .../parsers/legacy_loaders}/pdf_loaders/docling.py | 0 .../parsers/legacy_loaders}/pdf_loaders/docling2.py | 0 .../parsers/legacy_loaders}/pdf_loaders/dotsocr.py | 0 .../parsers/legacy_loaders}/pdf_loaders/marker.py | 0 .../parsers/legacy_loaders}/pdf_loaders/openai.py | 0 .../parsers/legacy_loaders}/pdf_loaders/pymupdf.py | 0 .../workers/parsers/legacy_loaders}/pptx_loader.py | 0 .../parsers/legacy_loaders}/test_base_loader.py | 0 .../parsers/legacy_loaders}/test_customdocloader.py | 2 +- .../parsers/legacy_loaders}/test_doc_loader.py | 8 ++++---- .../parsers/legacy_loaders}/test_docx_loader.py | 0 .../parsers/legacy_loaders}/test_eml_recursion.py | 4 ++-- .../workers/parsers/legacy_loaders}/txt_loader.py | 2 +- 31 files changed, 21 insertions(+), 29 deletions(-) delete mode 100644 openrag/components/indexer/loaders/serializer.py rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/CustomDocLoader.py (100%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/CustomHTMLLoader.py (100%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/__init__.py (97%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/audio/__init__.py (100%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/audio/local_whisper.py (96%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/audio/openai.py (100%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/audio/test_openai.py (100%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/base.py (100%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/doc.py (100%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/docx.py (100%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/eml_loader.py (100%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/image.py (100%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/pdf_loaders/__init__.py (100%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/pdf_loaders/docling.py (100%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/pdf_loaders/docling2.py (100%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/pdf_loaders/dotsocr.py (100%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/pdf_loaders/marker.py (100%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/pdf_loaders/openai.py (100%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/pdf_loaders/pymupdf.py (100%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/pptx_loader.py (100%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/test_base_loader.py (100%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/test_customdocloader.py (94%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/test_doc_loader.py (94%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/test_docx_loader.py (100%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/test_eml_recursion.py (95%) rename openrag/{components/indexer/loaders => services/workers/parsers/legacy_loaders}/txt_loader.py (97%) diff --git a/openrag/components/indexer/loaders/serializer.py b/openrag/components/indexer/loaders/serializer.py deleted file mode 100644 index 0f9c5d6c5..000000000 --- a/openrag/components/indexer/loaders/serializer.py +++ /dev/null @@ -1,11 +0,0 @@ -"""DocSerializer Ray actor — legacy re-export shim. - -The implementation now lives in -``services/workers/parsers/doc_serializer.py``; this module re-exports -``DocSerializer`` so existing import paths (``services/workers/bootstrap.py``, -``services/storage/serializer_ray_shim.py``) are unaffected. -""" - -from services.workers.parsers.doc_serializer import DocSerializer # noqa: F401 - -__all__ = ["DocSerializer"] diff --git a/openrag/services/inference/parsers/test_openai_audio.py b/openrag/services/inference/parsers/test_openai_audio.py index e6c665e2c..75c200e51 100644 --- a/openrag/services/inference/parsers/test_openai_audio.py +++ b/openrag/services/inference/parsers/test_openai_audio.py @@ -19,9 +19,12 @@ # ---- shim pydub before importing openai_audio ------------------------------ if "pydub" not in sys.modules: - pydub = types.ModuleType("pydub") - pydub.AudioSegment = MagicMock() # type: ignore[attr-defined] - sys.modules["pydub"] = pydub + try: + import pydub # noqa: E402,F401 — prefer the real library when it imports cleanly + except Exception: + fake_pydub = types.ModuleType("pydub") + fake_pydub.AudioSegment = MagicMock() # type: ignore[attr-defined] + sys.modules["pydub"] = fake_pydub from core.models.document import Document, DocumentType # noqa: E402 diff --git a/openrag/services/workers/bootstrap.py b/openrag/services/workers/bootstrap.py index 834d5dde3..89a81ef33 100644 --- a/openrag/services/workers/bootstrap.py +++ b/openrag/services/workers/bootstrap.py @@ -66,14 +66,14 @@ def get_task_state_manager(): def get_serializer(): - from components.indexer.loaders.serializer import DocSerializer + from services.workers.parsers.doc_serializer import DocSerializer return get_or_create_actor("DocSerializer", DocSerializer, lifetime="detached") def get_marker_pool(): - from components.indexer.loaders.pdf_loaders.docling2 import DoclingPool - from components.indexer.loaders.pdf_loaders.marker import MarkerPool + from services.workers.parsers.docling_workers import DoclingPool + from services.workers.parsers.marker_workers import MarkerPool config = _require_settings() pdf_loader = config.loader.file_loaders.pdf diff --git a/openrag/services/workers/parsers/doc_serializer.py b/openrag/services/workers/parsers/doc_serializer.py index 736d11a48..af53e433d 100644 --- a/openrag/services/workers/parsers/doc_serializer.py +++ b/openrag/services/workers/parsers/doc_serializer.py @@ -11,8 +11,8 @@ import ray import torch -from components.indexer.loaders import get_loader_classes from langchain_core.documents.base import Document +from services.workers.parsers.legacy_loaders import get_loader_classes @ray.remote(max_restarts=5) diff --git a/openrag/services/workers/parsers/doc_serializer_bridge.py b/openrag/services/workers/parsers/doc_serializer_bridge.py index 2c2f556b7..87b2e924e 100644 --- a/openrag/services/workers/parsers/doc_serializer_bridge.py +++ b/openrag/services/workers/parsers/doc_serializer_bridge.py @@ -12,7 +12,7 @@ class DocSerializerBridgeParser(DocumentParser): """Transitional parser backed by the legacy loader registry.""" def __init__(self, config: Any) -> None: - from components.indexer.loaders import get_loader_classes + from services.workers.parsers.legacy_loaders import get_loader_classes self._config = config self._loader_classes = get_loader_classes(config=config) diff --git a/openrag/components/indexer/loaders/CustomDocLoader.py b/openrag/services/workers/parsers/legacy_loaders/CustomDocLoader.py similarity index 100% rename from openrag/components/indexer/loaders/CustomDocLoader.py rename to openrag/services/workers/parsers/legacy_loaders/CustomDocLoader.py diff --git a/openrag/components/indexer/loaders/CustomHTMLLoader.py b/openrag/services/workers/parsers/legacy_loaders/CustomHTMLLoader.py similarity index 100% rename from openrag/components/indexer/loaders/CustomHTMLLoader.py rename to openrag/services/workers/parsers/legacy_loaders/CustomHTMLLoader.py diff --git a/openrag/components/indexer/loaders/__init__.py b/openrag/services/workers/parsers/legacy_loaders/__init__.py similarity index 97% rename from openrag/components/indexer/loaders/__init__.py rename to openrag/services/workers/parsers/legacy_loaders/__init__.py index 7f2688e34..b30b10575 100644 --- a/openrag/components/indexer/loaders/__init__.py +++ b/openrag/services/workers/parsers/legacy_loaders/__init__.py @@ -17,7 +17,7 @@ def get_loader_classes(config) -> dict[str, type[BaseLoader]]: # 1. Discover all subclasses - root_pkg = "components.indexer.loaders" + root_pkg = "services.workers.parsers.legacy_loaders" root_path = Path(__file__).parent discovered: dict[str, type[BaseLoader]] = {} diff --git a/openrag/components/indexer/loaders/audio/__init__.py b/openrag/services/workers/parsers/legacy_loaders/audio/__init__.py similarity index 100% rename from openrag/components/indexer/loaders/audio/__init__.py rename to openrag/services/workers/parsers/legacy_loaders/audio/__init__.py diff --git a/openrag/components/indexer/loaders/audio/local_whisper.py b/openrag/services/workers/parsers/legacy_loaders/audio/local_whisper.py similarity index 96% rename from openrag/components/indexer/loaders/audio/local_whisper.py rename to openrag/services/workers/parsers/legacy_loaders/audio/local_whisper.py index caa5cb3b1..51f9c8bc0 100644 --- a/openrag/components/indexer/loaders/audio/local_whisper.py +++ b/openrag/services/workers/parsers/legacy_loaders/audio/local_whisper.py @@ -6,7 +6,7 @@ implementation now live in ``services/workers/parsers/whisper_workers.py``; this module re-exports ``WhisperActor`` and ``WhisperPool`` for legacy import paths -(``components.indexer.loaders.audio.local_whisper.WhisperActor`` is +(``services.workers.parsers.legacy_loaders.audio.local_whisper.WhisperActor`` is still used by the OpenAI audio loader for language detection, and by ``services/workers/bootstrap.py`` for the actor bootstrap). diff --git a/openrag/components/indexer/loaders/audio/openai.py b/openrag/services/workers/parsers/legacy_loaders/audio/openai.py similarity index 100% rename from openrag/components/indexer/loaders/audio/openai.py rename to openrag/services/workers/parsers/legacy_loaders/audio/openai.py diff --git a/openrag/components/indexer/loaders/audio/test_openai.py b/openrag/services/workers/parsers/legacy_loaders/audio/test_openai.py similarity index 100% rename from openrag/components/indexer/loaders/audio/test_openai.py rename to openrag/services/workers/parsers/legacy_loaders/audio/test_openai.py diff --git a/openrag/components/indexer/loaders/base.py b/openrag/services/workers/parsers/legacy_loaders/base.py similarity index 100% rename from openrag/components/indexer/loaders/base.py rename to openrag/services/workers/parsers/legacy_loaders/base.py diff --git a/openrag/components/indexer/loaders/doc.py b/openrag/services/workers/parsers/legacy_loaders/doc.py similarity index 100% rename from openrag/components/indexer/loaders/doc.py rename to openrag/services/workers/parsers/legacy_loaders/doc.py diff --git a/openrag/components/indexer/loaders/docx.py b/openrag/services/workers/parsers/legacy_loaders/docx.py similarity index 100% rename from openrag/components/indexer/loaders/docx.py rename to openrag/services/workers/parsers/legacy_loaders/docx.py diff --git a/openrag/components/indexer/loaders/eml_loader.py b/openrag/services/workers/parsers/legacy_loaders/eml_loader.py similarity index 100% rename from openrag/components/indexer/loaders/eml_loader.py rename to openrag/services/workers/parsers/legacy_loaders/eml_loader.py diff --git a/openrag/components/indexer/loaders/image.py b/openrag/services/workers/parsers/legacy_loaders/image.py similarity index 100% rename from openrag/components/indexer/loaders/image.py rename to openrag/services/workers/parsers/legacy_loaders/image.py diff --git a/openrag/components/indexer/loaders/pdf_loaders/__init__.py b/openrag/services/workers/parsers/legacy_loaders/pdf_loaders/__init__.py similarity index 100% rename from openrag/components/indexer/loaders/pdf_loaders/__init__.py rename to openrag/services/workers/parsers/legacy_loaders/pdf_loaders/__init__.py diff --git a/openrag/components/indexer/loaders/pdf_loaders/docling.py b/openrag/services/workers/parsers/legacy_loaders/pdf_loaders/docling.py similarity index 100% rename from openrag/components/indexer/loaders/pdf_loaders/docling.py rename to openrag/services/workers/parsers/legacy_loaders/pdf_loaders/docling.py diff --git a/openrag/components/indexer/loaders/pdf_loaders/docling2.py b/openrag/services/workers/parsers/legacy_loaders/pdf_loaders/docling2.py similarity index 100% rename from openrag/components/indexer/loaders/pdf_loaders/docling2.py rename to openrag/services/workers/parsers/legacy_loaders/pdf_loaders/docling2.py diff --git a/openrag/components/indexer/loaders/pdf_loaders/dotsocr.py b/openrag/services/workers/parsers/legacy_loaders/pdf_loaders/dotsocr.py similarity index 100% rename from openrag/components/indexer/loaders/pdf_loaders/dotsocr.py rename to openrag/services/workers/parsers/legacy_loaders/pdf_loaders/dotsocr.py diff --git a/openrag/components/indexer/loaders/pdf_loaders/marker.py b/openrag/services/workers/parsers/legacy_loaders/pdf_loaders/marker.py similarity index 100% rename from openrag/components/indexer/loaders/pdf_loaders/marker.py rename to openrag/services/workers/parsers/legacy_loaders/pdf_loaders/marker.py diff --git a/openrag/components/indexer/loaders/pdf_loaders/openai.py b/openrag/services/workers/parsers/legacy_loaders/pdf_loaders/openai.py similarity index 100% rename from openrag/components/indexer/loaders/pdf_loaders/openai.py rename to openrag/services/workers/parsers/legacy_loaders/pdf_loaders/openai.py diff --git a/openrag/components/indexer/loaders/pdf_loaders/pymupdf.py b/openrag/services/workers/parsers/legacy_loaders/pdf_loaders/pymupdf.py similarity index 100% rename from openrag/components/indexer/loaders/pdf_loaders/pymupdf.py rename to openrag/services/workers/parsers/legacy_loaders/pdf_loaders/pymupdf.py diff --git a/openrag/components/indexer/loaders/pptx_loader.py b/openrag/services/workers/parsers/legacy_loaders/pptx_loader.py similarity index 100% rename from openrag/components/indexer/loaders/pptx_loader.py rename to openrag/services/workers/parsers/legacy_loaders/pptx_loader.py diff --git a/openrag/components/indexer/loaders/test_base_loader.py b/openrag/services/workers/parsers/legacy_loaders/test_base_loader.py similarity index 100% rename from openrag/components/indexer/loaders/test_base_loader.py rename to openrag/services/workers/parsers/legacy_loaders/test_base_loader.py diff --git a/openrag/components/indexer/loaders/test_customdocloader.py b/openrag/services/workers/parsers/legacy_loaders/test_customdocloader.py similarity index 94% rename from openrag/components/indexer/loaders/test_customdocloader.py rename to openrag/services/workers/parsers/legacy_loaders/test_customdocloader.py index c5c6d778e..2bd6f82f9 100644 --- a/openrag/components/indexer/loaders/test_customdocloader.py +++ b/openrag/services/workers/parsers/legacy_loaders/test_customdocloader.py @@ -13,7 +13,7 @@ @pytest.mark.asyncio async def test_customdocloader_accumulates_all_pages(tmp_path): - from components.indexer.loaders.CustomDocLoader import CustomDocLoader + from services.workers.parsers.legacy_loaders.CustomDocLoader import CustomDocLoader fake_pages = [ LCDocument(page_content="page-one"), diff --git a/openrag/components/indexer/loaders/test_doc_loader.py b/openrag/services/workers/parsers/legacy_loaders/test_doc_loader.py similarity index 94% rename from openrag/components/indexer/loaders/test_doc_loader.py rename to openrag/services/workers/parsers/legacy_loaders/test_doc_loader.py index a74e01b24..25e84f50a 100644 --- a/openrag/components/indexer/loaders/test_doc_loader.py +++ b/openrag/services/workers/parsers/legacy_loaders/test_doc_loader.py @@ -32,9 +32,9 @@ def metadata(): _PATCHES = [ - patch("components.indexer.loaders.doc.DocParser"), - patch("components.indexer.loaders.base.ChatOpenAI"), - patch("components.indexer.loaders.base.load_config"), + patch("services.workers.parsers.legacy_loaders.doc.DocParser"), + patch("services.workers.parsers.legacy_loaders.base.ChatOpenAI"), + patch("services.workers.parsers.legacy_loaders.base.load_config"), ] @@ -59,7 +59,7 @@ def _patch_cleanup(): def _make_loader(mock_config): - from components.indexer.loaders.doc import DocLoader + from services.workers.parsers.legacy_loaders.doc import DocLoader return DocLoader(config=mock_config) diff --git a/openrag/components/indexer/loaders/test_docx_loader.py b/openrag/services/workers/parsers/legacy_loaders/test_docx_loader.py similarity index 100% rename from openrag/components/indexer/loaders/test_docx_loader.py rename to openrag/services/workers/parsers/legacy_loaders/test_docx_loader.py diff --git a/openrag/components/indexer/loaders/test_eml_recursion.py b/openrag/services/workers/parsers/legacy_loaders/test_eml_recursion.py similarity index 95% rename from openrag/components/indexer/loaders/test_eml_recursion.py rename to openrag/services/workers/parsers/legacy_loaders/test_eml_recursion.py index e331c4245..b1644096b 100644 --- a/openrag/components/indexer/loaders/test_eml_recursion.py +++ b/openrag/services/workers/parsers/legacy_loaders/test_eml_recursion.py @@ -45,7 +45,7 @@ def _make_leaf_eml() -> bytes: @pytest.mark.asyncio async def test_eml_recursion_caps_at_max_depth(tmp_path): - from components.indexer.loaders.eml_loader import EmlLoader + from services.workers.parsers.legacy_loaders.eml_loader import EmlLoader # A single outer .eml whose attachment is another .eml. We seed the # call at depth = cap - 1, so processing the outer's attachment (which @@ -69,7 +69,7 @@ async def test_eml_recursion_caps_at_max_depth(tmp_path): @pytest.mark.asyncio async def test_eml_below_cap_does_not_skip(tmp_path): """At depth 0 the guard does not fire — attachments are still attempted.""" - from components.indexer.loaders.eml_loader import EmlLoader + from services.workers.parsers.legacy_loaders.eml_loader import EmlLoader eml_path = tmp_path / "nested.eml" eml_path.write_bytes(_make_eml_with_attached_eml(_make_leaf_eml())) diff --git a/openrag/components/indexer/loaders/txt_loader.py b/openrag/services/workers/parsers/legacy_loaders/txt_loader.py similarity index 97% rename from openrag/components/indexer/loaders/txt_loader.py rename to openrag/services/workers/parsers/legacy_loaders/txt_loader.py index 9e5154c10..82c4062c7 100644 --- a/openrag/components/indexer/loaders/txt_loader.py +++ b/openrag/services/workers/parsers/legacy_loaders/txt_loader.py @@ -14,13 +14,13 @@ import asyncio from pathlib import Path -from components.indexer.loaders.base import BaseLoader from core.indexing.parsers.markdown_parser import MarkdownParser from core.indexing.parsers.text_parser import TextParser from core.models.document import Document as CoreDocument from core.models.document import DocumentType from core.utils.logging import get_logger from langchain_core.documents.base import Document +from services.workers.parsers.legacy_loaders.base import BaseLoader logger = get_logger() From d391b63747603a7730ad07717ead0d12cf5a67af Mon Sep 17 00:00:00 2001 From: cbizeul Date: Fri, 29 May 2026 14:11:39 +0000 Subject: [PATCH 08/11] refactor: decouple api auth from services to satisfy layer guard MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The 12B import switch repointed two api files at services.auth directly, which the layer guard forbids (api -> di, core only; never services). Resolve it properly rather than bridging through di: - StateCookieSerializer / StateCookiePayload are pure (dataclasses + itsdangerous, no I/O), so move them to core/auth/state_cookie.py. api/routers/auth/oidc.py now imports from core (api -> core, allowed); services/auth re-exports them so existing services consumers are unaffected; the components.auth shim repoints to core. - refresh_session_if_needed is infrastructure and stays in services, but the middleware reached it via a direct module import. Expose it as AuthService.refresh_session_if_needed (a thin seam over the helper, passing self) so the middleware calls it on the AuthService it already obtains from di — consistent with every other session op it performs. Update the middleware test: the mock auth-service gains a default refresh delegating to the real helper, and the two refresh-outcome tests override that method instead of patching the module name. `scripts/check_layer_imports.py` is now clean (0 violations). --- openrag/api/middleware/auth.py | 7 +-- openrag/api/routers/auth/oidc.py | 2 +- openrag/components/auth/state_cookie.py | 4 +- openrag/components/auth/test_middleware.py | 43 +++++++++++-------- openrag/core/auth/__init__.py | 1 + .../{services => core}/auth/state_cookie.py | 0 openrag/services/auth/__init__.py | 3 +- .../services/orchestrators/auth_service.py | 15 +++++++ 8 files changed, 47 insertions(+), 28 deletions(-) create mode 100644 openrag/core/auth/__init__.py rename openrag/{services => core}/auth/state_cookie.py (100%) diff --git a/openrag/api/middleware/auth.py b/openrag/api/middleware/auth.py index fdb9adff4..6d5deaa8d 100644 --- a/openrag/api/middleware/auth.py +++ b/openrag/api/middleware/auth.py @@ -38,7 +38,6 @@ from core.utils.logging import get_logger from fastapi import Request from fastapi.responses import JSONResponse, RedirectResponse -from services.auth.refresh import refresh_session_if_needed from starlette.middleware.base import BaseHTTPMiddleware logger = get_logger() @@ -164,10 +163,9 @@ async def dispatch(self, request: Request, call_next): if cookie_token: session = await auth_service.get_oidc_session_by_token_for_request(cookie_token) if session is not None: - refreshed = await refresh_session_if_needed( + refreshed = await auth_service.refresh_session_if_needed( session=session, enc_key=enc_key, - auth_service=auth_service, ) if refreshed is None: # Refresh failed or session unusable → revoke and fall through. @@ -199,10 +197,9 @@ async def dispatch(self, request: Request, call_next): if auth_mode == "oidc": session = await auth_service.get_oidc_session_by_token_for_request(token) if session is not None: - refreshed = await refresh_session_if_needed( + refreshed = await auth_service.refresh_session_if_needed( session=session, enc_key=enc_key, - auth_service=auth_service, ) if refreshed is None: try: diff --git a/openrag/api/routers/auth/oidc.py b/openrag/api/routers/auth/oidc.py index ec5e91dc6..ddae2109f 100644 --- a/openrag/api/routers/auth/oidc.py +++ b/openrag/api/routers/auth/oidc.py @@ -23,12 +23,12 @@ import os +from core.auth.state_cookie import StateCookieSerializer from core.utils.exceptions import OpenRAGError from core.utils.logging import get_logger from di.providers import get_auth_service from fastapi import APIRouter, Depends, Form, HTTPException, Request, Response, status from fastapi.responses import JSONResponse, RedirectResponse -from services.auth.state_cookie import StateCookieSerializer logger = get_logger() router = APIRouter() diff --git a/openrag/components/auth/state_cookie.py b/openrag/components/auth/state_cookie.py index 0d00a0f52..647612541 100644 --- a/openrag/components/auth/state_cookie.py +++ b/openrag/components/auth/state_cookie.py @@ -1,5 +1,5 @@ -"""Re-export shim — implementation lives in services.auth.state_cookie.""" +"""Re-export shim — implementation lives in core.auth.state_cookie.""" -from services.auth.state_cookie import StateCookiePayload, StateCookieSerializer +from core.auth.state_cookie import StateCookiePayload, StateCookieSerializer __all__ = ["StateCookieSerializer", "StateCookiePayload"] diff --git a/openrag/components/auth/test_middleware.py b/openrag/components/auth/test_middleware.py index 622aeaa26..8fe329e57 100644 --- a/openrag/components/auth/test_middleware.py +++ b/openrag/components/auth/test_middleware.py @@ -47,6 +47,17 @@ def _make_auth_service_mock( mock.update_oidc_session_tokens_for_request = AsyncMock(return_value=None) + # Default refresh delegates to the real helper with this mock as the + # auth-service seam — mirrors the production middleware, which now reaches + # refresh through ``auth_service.refresh_session_if_needed`` rather than a + # module-level import. Tests needing a specific refresh outcome override it. + from services.auth.refresh import refresh_session_if_needed as _real_refresh + + async def _default_refresh(*, session, enc_key): + return await _real_refresh(session=session, enc_key=enc_key, auth_service=mock) + + mock.refresh_session_if_needed = AsyncMock(side_effect=_default_refresh) + return mock @@ -230,9 +241,9 @@ def test_cookie_near_expiry_triggers_refresh(self, monkeypatch): # Patch the helper at its import site inside the middleware module # to avoid any dependency on a real OIDC client. - async def fake_refresh(*, session, enc_key, auth_service): + async def fake_refresh(*, session, enc_key): new_exp = datetime.now() + timedelta(minutes=30) - await auth_service.update_oidc_session_tokens_for_request( + await vdb.update_oidc_session_tokens_for_request( session_id=session["id"], access_token_encrypted=b"new-enc-access", refresh_token_encrypted=b"new-enc-refresh", @@ -245,14 +256,11 @@ async def fake_refresh(*, session, enc_key, auth_service): "refresh_token_encrypted": b"new-enc-refresh", } - with patch( - "api.middleware.auth.refresh_session_if_needed", - side_effect=fake_refresh, - ): - app = _build_app(vdb) - with TestClient(app) as client: - client.cookies.set("openrag_session", "plain-cookie") - r = client.get("/v1/chat/completions") + vdb.refresh_session_if_needed = AsyncMock(side_effect=fake_refresh) + app = _build_app(vdb) + with TestClient(app) as client: + client.cookies.set("openrag_session", "plain-cookie") + r = client.get("/v1/chat/completions") assert r.status_code == 200 vdb.update_oidc_session_tokens_for_request.assert_awaited() @@ -263,17 +271,14 @@ def test_cookie_refresh_fails_session_revoked_and_302(self): session["access_token_expires_at"] = datetime.now() - timedelta(minutes=1) vdb = _make_auth_service_mock(user=None, session=session) - async def fake_refresh(*, session, enc_key, auth_service): + async def fake_refresh(*, session, enc_key): return None # refresh failed → invalid session - with patch( - "api.middleware.auth.refresh_session_if_needed", - side_effect=fake_refresh, - ): - app = _build_app(vdb) - with TestClient(app) as client: - client.cookies.set("openrag_session", "plain-cookie") - r = client.get("/", follow_redirects=False) + vdb.refresh_session_if_needed = AsyncMock(side_effect=fake_refresh) + app = _build_app(vdb) + with TestClient(app) as client: + client.cookies.set("openrag_session", "plain-cookie") + r = client.get("/", follow_redirects=False) assert r.status_code == 302 assert r.headers["location"].startswith("/auth/login?next=") diff --git a/openrag/core/auth/__init__.py b/openrag/core/auth/__init__.py new file mode 100644 index 000000000..adb70f698 --- /dev/null +++ b/openrag/core/auth/__init__.py @@ -0,0 +1 @@ +"""Auth domain primitives (pure — no infrastructure).""" diff --git a/openrag/services/auth/state_cookie.py b/openrag/core/auth/state_cookie.py similarity index 100% rename from openrag/services/auth/state_cookie.py rename to openrag/core/auth/state_cookie.py diff --git a/openrag/services/auth/__init__.py b/openrag/services/auth/__init__.py index 519dafcb8..560dc7967 100644 --- a/openrag/services/auth/__init__.py +++ b/openrag/services/auth/__init__.py @@ -1,9 +1,10 @@ """Auth service layer - OIDC client, session tokens, state cookie, and deps.""" +from core.auth.state_cookie import StateCookiePayload, StateCookieSerializer + from .deps import get_oidc_client, reset_oidc_client from .oidc_client import LogoutTokenClaims, OIDCClient, TokenBundle from .session_tokens import decrypt_token, encrypt_token, hash_session_token, issue_session_token -from .state_cookie import StateCookiePayload, StateCookieSerializer __all__ = [ "OIDCClient", diff --git a/openrag/services/orchestrators/auth_service.py b/openrag/services/orchestrators/auth_service.py index 8922c4b1a..91f0b1f50 100644 --- a/openrag/services/orchestrators/auth_service.py +++ b/openrag/services/orchestrators/auth_service.py @@ -362,6 +362,21 @@ async def update_oidc_session_tokens_for_request( async def revoke_oidc_session_by_id_for_request(self, session_id: int) -> None: await self._oidc_session_repo.revoke_session(session_id) + async def refresh_session_if_needed( + self, *, session: dict[str, Any], enc_key: str + ) -> dict[str, Any] | None: + """Refresh the IdP access token when it is near expiry. + + Thin seam over :func:`services.auth.refresh.refresh_session_if_needed`, + passing ``self`` as the auth-service the helper calls back into. Keeping + it on the orchestrator lets the middleware reach refresh through the same + ``AuthService`` it already uses for every other session operation — + the API layer never imports the services helper directly. + """ + from services.auth.refresh import refresh_session_if_needed as _refresh + + return await _refresh(session=session, enc_key=enc_key, auth_service=self) + # ------------------------------------------------------------------ # Auth policy — pure helpers (no I/O) # ------------------------------------------------------------------ From 9553397c4252734624e7a97c6da538e6bd8beed3 Mon Sep 17 00:00:00 2001 From: hedhoud <74668966+hedhoud@users.noreply.github.com> Date: Fri, 29 May 2026 16:17:57 +0200 Subject: [PATCH 09/11] style: format auth service --- openrag/services/orchestrators/auth_service.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/openrag/services/orchestrators/auth_service.py b/openrag/services/orchestrators/auth_service.py index 91f0b1f50..6344abdbe 100644 --- a/openrag/services/orchestrators/auth_service.py +++ b/openrag/services/orchestrators/auth_service.py @@ -362,9 +362,7 @@ async def update_oidc_session_tokens_for_request( async def revoke_oidc_session_by_id_for_request(self, session_id: int) -> None: await self._oidc_session_repo.revoke_session(session_id) - async def refresh_session_if_needed( - self, *, session: dict[str, Any], enc_key: str - ) -> dict[str, Any] | None: + async def refresh_session_if_needed(self, *, session: dict[str, Any], enc_key: str) -> dict[str, Any] | None: """Refresh the IdP access token when it is near expiry. Thin seam over :func:`services.auth.refresh.refresh_session_if_needed`, From 9ef1e934e75f5edc3a2ce22225ac1e8f56a75154 Mon Sep 17 00:00:00 2001 From: hedhoud <74668966+hedhoud@users.noreply.github.com> Date: Fri, 29 May 2026 16:43:40 +0200 Subject: [PATCH 10/11] refactor: remove legacy package roots --- openrag/api/dependencies/files.py | 15 +- openrag/api/dependencies/test_files.py | 25 + openrag/api/main.py | 2 +- openrag/api/mcp/server.py | 2 +- openrag/api/routers/user/chat.py | 2 +- openrag/api/schemas/user/chat.py | 2 +- openrag/chainlit_api.py | 2 +- openrag/components/__init__.py | 1 - openrag/components/auth/__init__.py | 18 - openrag/components/auth/deps.py | 24 - openrag/components/auth/oidc_client.py | 5 - openrag/components/auth/refresh.py | 5 - openrag/components/auth/session_tokens.py | 5 - openrag/components/auth/state_cookie.py | 5 - openrag/components/auth/test_middleware.py | 515 -------------- openrag/components/auth/test_oidc_client.py | 433 ------------ .../components/auth/test_session_tokens.py | 68 -- openrag/components/auth/test_state_cookie.py | 67 -- openrag/components/files.py | 43 -- openrag/components/indexer/__init__.py | 8 - .../components/indexer/chunker/__init__.py | 1 - openrag/components/indexer/chunker/chunker.py | 222 ------ .../indexer/chunker/test_chunking.py | 232 ------ openrag/components/indexer/chunker/utils.py | 33 - .../components/indexer/embeddings/__init__.py | 23 - openrag/components/indexer/embeddings/base.py | 20 - .../components/indexer/embeddings/openai.py | 155 ---- openrag/components/indexer/utils/__init__.py | 2 - openrag/components/indexer/utils/files.py | 89 --- .../indexer/utils/test_text_sanitizer.py | 200 ------ .../indexer/utils/text_sanitizer.py | 7 - openrag/components/llm.py | 149 ---- openrag/components/prompts/__init__.py | 1 - openrag/components/prompts/prompts.py | 42 -- openrag/components/ray_utils.py | 6 - openrag/components/reranker.py | 82 --- openrag/components/reranker/__init__.py | 43 -- openrag/components/reranker/base.py | 18 - openrag/components/reranker/infinity.py | 63 -- openrag/components/reranker/openai.py | 64 -- .../components/reranker/test_rrf_reranking.py | 75 -- openrag/components/test_files.py | 55 -- openrag/components/test_llm.py | 97 --- openrag/components/test_relationships.py | 669 ------------------ openrag/components/utils.py | 326 --------- openrag/config/__init__.py | 39 - openrag/config/loader.py | 296 -------- openrag/config/models.py | 538 -------------- openrag/core/utils/singleton.py | 9 + openrag/core/utils/source_filtering.py | 12 +- openrag/core/utils/test_source_filtering.py | 12 + openrag/scripts/backup.py | 2 +- openrag/scripts/restore.py | 2 +- .../inference/parsers/test_openai_audio.py | 2 +- openrag/services/inference/runtime.py | 13 +- openrag/services/inference/test_runtime.py | 34 + openrag/services/persistence/document_repo.py | 2 +- .../persistence/migrations/alembic/env.py | 2 +- .../1.add_created_at_temporal_fields.py | 2 +- .../persistence/migrations/milvus/migrate.py | 2 +- .../test_ancestor_recursion_cap.py | 6 +- openrag/services/workers/indexer_pool.py | 4 +- .../workers/parsers/doc_serializer.py | 2 +- .../workers/parsers/doc_serializer_adapter.py | 2 +- .../workers/parsers/docling_workers.py | 2 +- .../workers/parsers/legacy_loaders/base.py | 14 +- .../legacy_loaders/pdf_loaders/docling.py | 2 +- .../parsers/legacy_loaders/test_doc_loader.py | 3 +- .../legacy_loaders/test_eml_recursion.py | 2 +- .../workers/parsers/marker_workers.py | 6 +- .../workers/parsers/whisper_workers.py | 6 +- openrag/services/workers/task_state.py | 2 +- openrag/services/workers/test_indexer_pool.py | 14 + .../tests/test_relationships_integration.py | 140 +--- openrag/utils/__init__.py | 0 openrag/utils/exceptions/__init__.py | 3 - openrag/utils/exceptions/base.py | 7 - openrag/utils/exceptions/embeddings.py | 7 - openrag/utils/exceptions/vectordb.py | 17 - openrag/utils/external_resource_errors.py | 13 - 80 files changed, 171 insertions(+), 4969 deletions(-) delete mode 100644 openrag/components/__init__.py delete mode 100644 openrag/components/auth/__init__.py delete mode 100644 openrag/components/auth/deps.py delete mode 100644 openrag/components/auth/oidc_client.py delete mode 100644 openrag/components/auth/refresh.py delete mode 100644 openrag/components/auth/session_tokens.py delete mode 100644 openrag/components/auth/state_cookie.py delete mode 100644 openrag/components/auth/test_middleware.py delete mode 100644 openrag/components/auth/test_oidc_client.py delete mode 100644 openrag/components/auth/test_session_tokens.py delete mode 100644 openrag/components/auth/test_state_cookie.py delete mode 100644 openrag/components/files.py delete mode 100644 openrag/components/indexer/__init__.py delete mode 100644 openrag/components/indexer/chunker/__init__.py delete mode 100644 openrag/components/indexer/chunker/chunker.py delete mode 100644 openrag/components/indexer/chunker/test_chunking.py delete mode 100644 openrag/components/indexer/chunker/utils.py delete mode 100644 openrag/components/indexer/embeddings/__init__.py delete mode 100644 openrag/components/indexer/embeddings/base.py delete mode 100644 openrag/components/indexer/embeddings/openai.py delete mode 100644 openrag/components/indexer/utils/__init__.py delete mode 100644 openrag/components/indexer/utils/files.py delete mode 100644 openrag/components/indexer/utils/test_text_sanitizer.py delete mode 100644 openrag/components/indexer/utils/text_sanitizer.py delete mode 100644 openrag/components/llm.py delete mode 100644 openrag/components/prompts/__init__.py delete mode 100644 openrag/components/prompts/prompts.py delete mode 100644 openrag/components/ray_utils.py delete mode 100644 openrag/components/reranker.py delete mode 100644 openrag/components/reranker/__init__.py delete mode 100644 openrag/components/reranker/base.py delete mode 100644 openrag/components/reranker/infinity.py delete mode 100644 openrag/components/reranker/openai.py delete mode 100644 openrag/components/reranker/test_rrf_reranking.py delete mode 100644 openrag/components/test_files.py delete mode 100644 openrag/components/test_llm.py delete mode 100644 openrag/components/test_relationships.py delete mode 100644 openrag/components/utils.py delete mode 100644 openrag/config/__init__.py delete mode 100644 openrag/config/loader.py delete mode 100644 openrag/config/models.py create mode 100644 openrag/core/utils/singleton.py create mode 100644 openrag/services/inference/test_runtime.py delete mode 100644 openrag/utils/__init__.py delete mode 100644 openrag/utils/exceptions/__init__.py delete mode 100644 openrag/utils/exceptions/base.py delete mode 100644 openrag/utils/exceptions/embeddings.py delete mode 100644 openrag/utils/exceptions/vectordb.py delete mode 100644 openrag/utils/external_resource_errors.py diff --git a/openrag/api/dependencies/files.py b/openrag/api/dependencies/files.py index 5f262a13b..ab97d0023 100644 --- a/openrag/api/dependencies/files.py +++ b/openrag/api/dependencies/files.py @@ -4,6 +4,7 @@ import aiofiles import consts from core.indexing import validators as core_validators +from core.utils.exceptions import ValidationError from core.utils.filename import make_unique_filename from di.providers import get_config from fastapi import Depends, Form, UploadFile @@ -43,9 +44,19 @@ async def save_file_to_disk( ) -> Path: """Save an uploaded file to disk in chunks and return the saved path.""" dest_dir.mkdir(parents=True, exist_ok=True) + dest_dir = dest_dir.resolve() - filename = make_unique_filename(file.filename) if with_random_prefix else file.filename - file_path = dest_dir / filename + raw_filename = file.filename or "" + safe_basename = raw_filename.replace("\\", "/").split("/")[-1].strip() + if not safe_basename: + raise ValidationError("Uploaded file must have a filename.", status_code=400) + + filename = make_unique_filename(safe_basename) if with_random_prefix else safe_basename + file_path = (dest_dir / filename).resolve() + try: + file_path.relative_to(dest_dir) + except ValueError: + raise ValidationError("Uploaded filename resolves outside destination directory.", status_code=400) async with aiofiles.open(file_path, "wb") as buffer: while True: diff --git a/openrag/api/dependencies/test_files.py b/openrag/api/dependencies/test_files.py index ab5b9e40b..1d930337a 100644 --- a/openrag/api/dependencies/test_files.py +++ b/openrag/api/dependencies/test_files.py @@ -60,6 +60,31 @@ def fake_make_unique_filename(filename: str) -> str: assert saved_path.read_bytes() == file_content +@pytest.mark.asyncio +async def test_save_file_to_disk_strips_path_components(tmp_path): + upload = UploadFile( + filename="../../nested/evil.txt", + file=io.BytesIO(b"content"), + ) + + saved_path = await save_file_to_disk(file=upload, dest_dir=tmp_path) + + assert saved_path.parent == tmp_path.resolve() + assert saved_path.name == "evil.txt" + assert saved_path.read_bytes() == b"content" + + +@pytest.mark.asyncio +async def test_save_file_to_disk_rejects_empty_filename(tmp_path): + upload = UploadFile( + filename="", + file=io.BytesIO(b"content"), + ) + + with pytest.raises(ValidationError): + await save_file_to_disk(file=upload, dest_dir=tmp_path) + + @pytest.mark.parametrize( "input_name,expected", [ diff --git a/openrag/api/main.py b/openrag/api/main.py index 89012cadf..d683e6e07 100644 --- a/openrag/api/main.py +++ b/openrag/api/main.py @@ -51,7 +51,7 @@ from api.routers.user.extract import router as extract_router from api.routers.user.health import router as health_router from api.routers.user.search import router as search_router -from config import load_config +from core.config import load_config from core.utils.logging import get_logger from di.container import ServiceContainer from di.providers import set_container diff --git a/openrag/api/mcp/server.py b/openrag/api/mcp/server.py index 4d5ca5c11..6df49989a 100644 --- a/openrag/api/mcp/server.py +++ b/openrag/api/mcp/server.py @@ -35,7 +35,7 @@ reset_auth_context, set_auth_context, ) -from config import load_config +from core.config import load_config from core.utils.log_tail import app_log_file from core.utils.logging import get_logger from di.container import ServiceContainer diff --git a/openrag/api/routers/user/chat.py b/openrag/api/routers/user/chat.py index 03badfc18..b156d7898 100644 --- a/openrag/api/routers/user/chat.py +++ b/openrag/api/routers/user/chat.py @@ -30,7 +30,7 @@ ) from api.routers.user.source_links import build_document_source_link from api.schemas.user.chat import OpenAIChatCompletionRequest, OpenAICompletionRequest -from config import load_config +from core.config import load_config from core.utils.exceptions import OpenRAGError from core.utils.logging import get_logger from core.utils.text import get_num_tokens, sanitize_text diff --git a/openrag/api/schemas/user/chat.py b/openrag/api/schemas/user/chat.py index f90e67a9e..0c3347c2a 100644 --- a/openrag/api/schemas/user/chat.py +++ b/openrag/api/schemas/user/chat.py @@ -4,7 +4,7 @@ def default_max_tokens(): - from config import load_config + from core.config import load_config return load_config().llm_context.max_output_tokens diff --git a/openrag/chainlit_api.py b/openrag/chainlit_api.py index dad479f32..61c3636e5 100644 --- a/openrag/chainlit_api.py +++ b/openrag/chainlit_api.py @@ -6,7 +6,7 @@ def _get_auth_service(request): container = getattr(request.app.state, "container", None) if container is None: - from config import load_config + from core.config import load_config from di.container import ServiceContainer container = ServiceContainer(load_config()) diff --git a/openrag/components/__init__.py b/openrag/components/__init__.py deleted file mode 100644 index 8b1378917..000000000 --- a/openrag/components/__init__.py +++ /dev/null @@ -1 +0,0 @@ - diff --git a/openrag/components/auth/__init__.py b/openrag/components/auth/__init__.py deleted file mode 100644 index a1eb42365..000000000 --- a/openrag/components/auth/__init__.py +++ /dev/null @@ -1,18 +0,0 @@ -from components.auth.deps import get_oidc_client, reset_oidc_client -from components.auth.oidc_client import LogoutTokenClaims, OIDCClient, TokenBundle -from components.auth.session_tokens import decrypt_token, encrypt_token, hash_session_token, issue_session_token -from components.auth.state_cookie import StateCookiePayload, StateCookieSerializer - -__all__ = [ - "OIDCClient", - "TokenBundle", - "LogoutTokenClaims", - "issue_session_token", - "encrypt_token", - "decrypt_token", - "hash_session_token", - "StateCookieSerializer", - "StateCookiePayload", - "get_oidc_client", - "reset_oidc_client", -] diff --git a/openrag/components/auth/deps.py b/openrag/components/auth/deps.py deleted file mode 100644 index 5038ca94b..000000000 --- a/openrag/components/auth/deps.py +++ /dev/null @@ -1,24 +0,0 @@ -"""Compatibility shim - implementation lives in services.auth.deps.""" - -from __future__ import annotations - -from services.auth.deps import get_oidc_client as _service_get_oidc_client -from services.auth.deps import reset_oidc_client as _service_reset_oidc_client -from services.auth.oidc_client import OIDCClient - -_client: OIDCClient | None = None - - -def get_oidc_client() -> OIDCClient: - if _client is not None: - return _client - return _service_get_oidc_client() - - -def reset_oidc_client() -> None: - global _client - _client = None - _service_reset_oidc_client() - - -__all__ = ["get_oidc_client", "reset_oidc_client"] diff --git a/openrag/components/auth/oidc_client.py b/openrag/components/auth/oidc_client.py deleted file mode 100644 index c71f542a5..000000000 --- a/openrag/components/auth/oidc_client.py +++ /dev/null @@ -1,5 +0,0 @@ -"""Re-export shim — implementation lives in services.auth.oidc_client.""" - -from services.auth.oidc_client import LogoutTokenClaims, OIDCClient, TokenBundle - -__all__ = ["OIDCClient", "TokenBundle", "LogoutTokenClaims"] diff --git a/openrag/components/auth/refresh.py b/openrag/components/auth/refresh.py deleted file mode 100644 index 75bcea848..000000000 --- a/openrag/components/auth/refresh.py +++ /dev/null @@ -1,5 +0,0 @@ -"""Re-export shim — implementation lives in services.auth.refresh.""" - -from services.auth.refresh import refresh_session_if_needed - -__all__ = ["refresh_session_if_needed"] diff --git a/openrag/components/auth/session_tokens.py b/openrag/components/auth/session_tokens.py deleted file mode 100644 index 48a228db8..000000000 --- a/openrag/components/auth/session_tokens.py +++ /dev/null @@ -1,5 +0,0 @@ -"""Re-export shim — implementation lives in services.auth.session_tokens.""" - -from services.auth.session_tokens import decrypt_token, encrypt_token, hash_session_token, issue_session_token - -__all__ = ["issue_session_token", "hash_session_token", "encrypt_token", "decrypt_token"] diff --git a/openrag/components/auth/state_cookie.py b/openrag/components/auth/state_cookie.py deleted file mode 100644 index 647612541..000000000 --- a/openrag/components/auth/state_cookie.py +++ /dev/null @@ -1,5 +0,0 @@ -"""Re-export shim — implementation lives in core.auth.state_cookie.""" - -from core.auth.state_cookie import StateCookiePayload, StateCookieSerializer - -__all__ = ["StateCookieSerializer", "StateCookiePayload"] diff --git a/openrag/components/auth/test_middleware.py b/openrag/components/auth/test_middleware.py deleted file mode 100644 index 8fe329e57..000000000 --- a/openrag/components/auth/test_middleware.py +++ /dev/null @@ -1,515 +0,0 @@ -"""Unit tests for the Phase-5 ``AuthMiddleware``. - -These tests mount the middleware on a minimal FastAPI app with a ``MagicMock`` -``vectordb`` — no Ray, no Postgres, no Milvus. They exercise the decision tree -documented in ``.omc/plans/oidc-auth/plan.md`` §6.1. - -Timezone policy: Phase 2 stores session timestamps as naive local time -(``datetime.now()``), so the refresh helper compares naive datetimes. These -tests follow suit. -""" - -from __future__ import annotations - -from datetime import datetime, timedelta -from unittest.mock import AsyncMock, MagicMock, patch - -import pytest -from api.middleware.auth import AuthMiddleware, is_ui_path -from fastapi import FastAPI, Request -from fastapi.testclient import TestClient - -# --------------------------------------------------------------------------- -# Helpers -# --------------------------------------------------------------------------- - - -def _make_auth_service_mock( - *, - user=None, - user_by_token=None, - session=None, - partitions=None, -): - """Return a MagicMock exposing the auth-service surface used by the middleware.""" - mock = MagicMock() - - mock.get_user_for_request = AsyncMock(return_value=user or {"id": 1, "display_name": "Admin"}) - - mock.get_user_by_token_for_request = AsyncMock(return_value=user_by_token) - - mock.get_oidc_session_by_token_for_request = AsyncMock(return_value=session) - mock.get_oidc_session_by_id_for_request = AsyncMock(return_value=session) - - mock.list_user_partitions_for_request = AsyncMock(return_value=partitions or []) - - mock.revoke_oidc_session_by_id_for_request = AsyncMock(return_value=None) - - mock.update_oidc_session_tokens_for_request = AsyncMock(return_value=None) - - # Default refresh delegates to the real helper with this mock as the - # auth-service seam — mirrors the production middleware, which now reaches - # refresh through ``auth_service.refresh_session_if_needed`` rather than a - # module-level import. Tests needing a specific refresh outcome override it. - from services.auth.refresh import refresh_session_if_needed as _real_refresh - - async def _default_refresh(*, session, enc_key): - return await _real_refresh(session=session, enc_key=enc_key, auth_service=mock) - - mock.refresh_session_if_needed = AsyncMock(side_effect=_default_refresh) - - return mock - - -def _build_app(auth_service_mock) -> FastAPI: - """Construct a FastAPI app with the middleware under test.""" - app = FastAPI() - app.add_middleware(AuthMiddleware, get_auth_service=lambda _request: auth_service_mock) - - @app.get("/") - async def root(request: Request): - return {"user": request.state.user["id"]} - - @app.get("/v1/chat/completions") - async def chat(request: Request): - return {"user": request.state.user["id"]} - - @app.get("/indexer/foo") - async def indexer_foo(request: Request): - return {"user": request.state.user["id"]} - - @app.get("/users/info") - async def users_info(request: Request): - return {"user": request.state.user["id"]} - - @app.get("/static/foo.pdf") - async def static_file(request: Request): - return {"user": request.state.user["id"]} - - @app.get("/health_check") - async def hc(): - return "ok" - - return app - - -# --------------------------------------------------------------------------- -# is_ui_path — pure function -# --------------------------------------------------------------------------- - - -@pytest.mark.parametrize( - "path,expected", - [ - ("/", True), - ("/static/x.pdf", True), - ("/static", True), - ("/v1/chat/completions", False), - ("/v1/models", False), - ("/indexer/add_file", False), - ("/search/foo", False), - ("/users/info", False), - ("/partition/foo", False), - ("/workspaces/list", False), - ("/queue/info", False), - ("/extract/something", False), - ("/actors/", False), - ("/monitoring/status", False), - ("/tools/execute", False), - ("/unknown/thing", False), # default: not UI (avoid redirect loops) - ], -) -def test_is_ui_path(path, expected): - assert is_ui_path(path) is expected - - -# --------------------------------------------------------------------------- -# Token mode (legacy) — must preserve 403 + legacy error bodies -# --------------------------------------------------------------------------- - - -class TestTokenModeLegacy: - @pytest.fixture(autouse=True) - def _env(self, monkeypatch): - monkeypatch.setenv("AUTH_MODE", "token") - monkeypatch.setenv("AUTH_TOKEN", "configured-admin-token") - - def test_bearer_valid_returns_200(self): - vdb = _make_auth_service_mock(user_by_token={"id": 7, "display_name": "U"}) - app = _build_app(vdb) - with TestClient(app) as client: - r = client.get( - "/v1/chat/completions", - headers={"Authorization": "Bearer good-token"}, - ) - assert r.status_code == 200 - assert r.json() == {"user": 7} - - def test_bearer_invalid_returns_403(self): - vdb = _make_auth_service_mock(user_by_token=None) - app = _build_app(vdb) - with TestClient(app) as client: - r = client.get( - "/v1/chat/completions", - headers={"Authorization": "Bearer bogus"}, - ) - assert r.status_code == 403 - assert r.json() == {"detail": "Invalid token"} - - def test_missing_token_returns_403(self): - vdb = _make_auth_service_mock() - app = _build_app(vdb) - with TestClient(app) as client: - r = client.get("/v1/chat/completions") - assert r.status_code == 403 - assert r.json() == {"detail": "Missing token"} - - def test_bypass_path_open(self): - vdb = _make_auth_service_mock() - app = _build_app(vdb) - with TestClient(app) as client: - r = client.get("/health_check") - assert r.status_code == 200 - - -class TestTokenModeDevBypass: - """AUTH_MODE=token, AUTH_TOKEN unset → all requests resolve to user id=1.""" - - def test_no_token_resolves_user_1(self, monkeypatch): - monkeypatch.setenv("AUTH_MODE", "token") - monkeypatch.delenv("AUTH_TOKEN", raising=False) - vdb = _make_auth_service_mock(user={"id": 1, "display_name": "Admin"}) - app = _build_app(vdb) - with TestClient(app) as client: - r = client.get("/v1/chat/completions") - assert r.status_code == 200 - assert r.json() == {"user": 1} - vdb.get_user_for_request.assert_awaited_with(1) - - -# --------------------------------------------------------------------------- -# OIDC mode -# --------------------------------------------------------------------------- - - -class TestOIDCMode: - @pytest.fixture(autouse=True) - def _env(self, monkeypatch): - monkeypatch.setenv("AUTH_MODE", "oidc") - # refresh helper reads OIDC_TOKEN_ENCRYPTION_KEY but we patch the helper - # in refresh-related tests, so a dummy value is fine. - monkeypatch.setenv("OIDC_TOKEN_ENCRYPTION_KEY", "dummy") - - def _fresh_session(self, user_id=42): - """A session whose access_token is still well within lifetime.""" - return { - "id": 1, - "user_id": user_id, - "sub": "sub-abc", - "sid": "sid-xyz", - "id_token_encrypted": None, - "access_token_encrypted": b"enc-access", - "refresh_token_encrypted": b"enc-refresh", - "access_token_expires_at": datetime.now() + timedelta(minutes=30), - "session_expires_at": datetime.now() + timedelta(hours=8), - "revoked_at": None, - "last_refresh_at": None, - } - - # -- cookie session happy path ------------------------------------------ - - def test_cookie_valid_and_access_token_fresh_no_refresh(self): - session = self._fresh_session(user_id=42) - user = {"id": 42, "display_name": "Alice"} - vdb = _make_auth_service_mock(user=user, session=session) - app = _build_app(vdb) - with TestClient(app) as client: - client.cookies.set("openrag_session", "plain-cookie") - r = client.get("/v1/chat/completions") - assert r.status_code == 200 - assert r.json() == {"user": 42} - vdb.update_oidc_session_tokens_for_request.assert_not_awaited() - vdb.revoke_oidc_session_by_id_for_request.assert_not_awaited() - - def test_cookie_near_expiry_triggers_refresh(self, monkeypatch): - """access_token within 60s of expiry AND refresh_token present → refresh.""" - session = self._fresh_session(user_id=42) - # Force the refresh helper to "see" the token as near-expiry. - session["access_token_expires_at"] = datetime.now() + timedelta(seconds=5) - user = {"id": 42} - vdb = _make_auth_service_mock(user=user, session=session) - - # Patch the helper at its import site inside the middleware module - # to avoid any dependency on a real OIDC client. - async def fake_refresh(*, session, enc_key): - new_exp = datetime.now() + timedelta(minutes=30) - await vdb.update_oidc_session_tokens_for_request( - session_id=session["id"], - access_token_encrypted=b"new-enc-access", - refresh_token_encrypted=b"new-enc-refresh", - access_token_expires_at=new_exp, - ) - return { - **session, - "access_token_encrypted": b"new-enc-access", - "access_token_expires_at": new_exp, - "refresh_token_encrypted": b"new-enc-refresh", - } - - vdb.refresh_session_if_needed = AsyncMock(side_effect=fake_refresh) - app = _build_app(vdb) - with TestClient(app) as client: - client.cookies.set("openrag_session", "plain-cookie") - r = client.get("/v1/chat/completions") - - assert r.status_code == 200 - vdb.update_oidc_session_tokens_for_request.assert_awaited() - - def test_cookie_refresh_fails_session_revoked_and_302(self): - """access_token expired + refresh fails → session revoked, UI request → 302.""" - session = self._fresh_session(user_id=42) - session["access_token_expires_at"] = datetime.now() - timedelta(minutes=1) - vdb = _make_auth_service_mock(user=None, session=session) - - async def fake_refresh(*, session, enc_key): - return None # refresh failed → invalid session - - vdb.refresh_session_if_needed = AsyncMock(side_effect=fake_refresh) - app = _build_app(vdb) - with TestClient(app) as client: - client.cookies.set("openrag_session", "plain-cookie") - r = client.get("/", follow_redirects=False) - - assert r.status_code == 302 - assert r.headers["location"].startswith("/auth/login?next=") - vdb.revoke_oidc_session_by_id_for_request.assert_awaited_with(1) - - # -- bearer fallback ---------------------------------------------------- - - def test_bearer_fallback_accepted_in_oidc_mode(self): - """Programmatic clients keep using ``users.token`` in oidc mode.""" - vdb = _make_auth_service_mock(user_by_token={"id": 9, "display_name": "bot"}) - app = _build_app(vdb) - with TestClient(app) as client: - r = client.get( - "/v1/chat/completions", - headers={"Authorization": "Bearer ci-token"}, - ) - assert r.status_code == 200 - assert r.json() == {"user": 9} - - # -- unauthenticated branching ------------------------------------------ - - def test_no_creds_api_path_returns_401(self): - vdb = _make_auth_service_mock() - app = _build_app(vdb) - with TestClient(app) as client: - r = client.get("/indexer/foo") - assert r.status_code == 401 - assert r.json() == {"detail": "Unauthenticated"} - - def test_no_creds_root_path_returns_302(self): - vdb = _make_auth_service_mock() - app = _build_app(vdb) - with TestClient(app) as client: - r = client.get("/", follow_redirects=False) - assert r.status_code == 302 - # ``next`` must preserve the original path+query - assert r.headers["location"] == "/auth/login?next=%2F" - - def test_no_creds_root_with_query_preserves_next(self): - vdb = _make_auth_service_mock() - app = _build_app(vdb) - with TestClient(app) as client: - r = client.get("/?foo=bar", follow_redirects=False) - assert r.status_code == 302 - # %2F / %3F / %3D — full url-encoding - assert "next=" in r.headers["location"] - assert "%2F" in r.headers["location"] - - def test_no_creds_static_path_returns_302(self): - vdb = _make_auth_service_mock() - app = _build_app(vdb) - with TestClient(app) as client: - r = client.get("/static/foo.pdf", follow_redirects=False) - assert r.status_code == 302 - - def test_no_creds_v1_chat_returns_401(self): - vdb = _make_auth_service_mock() - app = _build_app(vdb) - with TestClient(app) as client: - r = client.get("/v1/chat/completions") - assert r.status_code == 401 - - -# --------------------------------------------------------------------------- -# refresh_session_if_needed — behavioural unit test -# --------------------------------------------------------------------------- - - -class TestRefreshHelper: - @pytest.mark.asyncio - async def test_no_refresh_when_token_fresh(self): - from components.auth.refresh import refresh_session_if_needed - - session = { - "id": 1, - "access_token_expires_at": datetime.now() + timedelta(minutes=30), - "refresh_token_encrypted": b"foo", - } - vdb = MagicMock() - vdb.update_oidc_session_tokens_for_request = MagicMock() - vdb.update_oidc_session_tokens_for_request = AsyncMock() - - out = await refresh_session_if_needed(session=session, enc_key="k", auth_service=vdb) - assert out is session - vdb.update_oidc_session_tokens_for_request.assert_not_awaited() - - @pytest.mark.asyncio - async def test_expired_no_refresh_token_returns_none(self): - from components.auth.refresh import refresh_session_if_needed - - session = { - "id": 1, - "access_token_expires_at": datetime.now() - timedelta(minutes=1), - "refresh_token_encrypted": None, - } - vdb = MagicMock() - out = await refresh_session_if_needed(session=session, enc_key="k", auth_service=vdb) - assert out is None - - # ------------------------------------------------------------------ - # M1: refresh-token stampede guard - # ------------------------------------------------------------------ - - @pytest.mark.asyncio - async def test_refresh_short_circuit_when_last_refresh_recent(self): - """If another request refreshed <5s ago, reuse the fresh row; do NOT - hit the IdP again with a refresh_token that has already been rotated.""" - from services.auth import refresh as refresh_mod - from services.auth.refresh import refresh_session_if_needed - - now = datetime.now() - fresh_exp = now + timedelta(minutes=30) - fresh_row = { - "id": 1, - "access_token_expires_at": fresh_exp, - "refresh_token_encrypted": b"new-refresh", - "access_token_encrypted": b"new-access", - "last_refresh_at": now, - } - stale_session = { - "id": 1, - # About to expire → normally we would call the IdP. - "access_token_expires_at": now + timedelta(seconds=5), - "refresh_token_encrypted": b"old-refresh", - "last_refresh_at": now - timedelta(seconds=2), # sibling just refreshed - } - - vdb = MagicMock() - vdb.get_oidc_session_by_id_for_request = MagicMock() - vdb.get_oidc_session_by_id_for_request = AsyncMock(return_value=fresh_row) - vdb.update_oidc_session_tokens_for_request = MagicMock() - vdb.update_oidc_session_tokens_for_request = AsyncMock() - - # Sentinel: the IdP client must NOT be contacted. - fake_client = MagicMock() - fake_client.refresh_access_token = AsyncMock( - side_effect=AssertionError("IdP must not be called during stampede short-circuit") - ) - with patch.object(refresh_mod, "get_oidc_client", return_value=fake_client): - out = await refresh_session_if_needed(session=stale_session, enc_key="k", auth_service=vdb) - - assert out is fresh_row - fake_client.refresh_access_token.assert_not_awaited() - vdb.update_oidc_session_tokens_for_request.assert_not_awaited() - - @pytest.mark.asyncio - async def test_refresh_recovers_when_idp_rejects_stale_refresh_token(self): - """IdP rejects our refresh_token (sibling already rotated it); the helper - re-reads the session and returns the sibling's fresh tokens.""" - from services.auth import refresh as refresh_mod - from services.auth.refresh import refresh_session_if_needed - - now = datetime.now() - stale_session = { - "id": 1, - "access_token_expires_at": now + timedelta(seconds=5), - "refresh_token_encrypted": b"old-refresh", - # No recent last_refresh_at → stampede short-circuit does NOT fire. - "last_refresh_at": None, - } - fresh_row = { - "id": 1, - "access_token_expires_at": now + timedelta(minutes=30), - "refresh_token_encrypted": b"new-refresh", - "access_token_encrypted": b"new-access", - "last_refresh_at": now, - } - - vdb = MagicMock() - vdb.get_oidc_session_by_id_for_request = MagicMock() - vdb.get_oidc_session_by_id_for_request = AsyncMock(return_value=fresh_row) - - fake_client = MagicMock() - fake_client.refresh_access_token = AsyncMock(side_effect=RuntimeError("invalid_grant")) - with ( - patch.object(refresh_mod, "get_oidc_client", return_value=fake_client), - patch.object(refresh_mod, "decrypt_token", return_value="old-refresh-plain"), - ): - out = await refresh_session_if_needed(session=stale_session, enc_key="k", auth_service=vdb) - - assert out is fresh_row - fake_client.refresh_access_token.assert_awaited_once() - vdb.get_oidc_session_by_id_for_request.assert_awaited_once_with(1) - - @pytest.mark.asyncio - async def test_refresh_returns_none_when_idp_rejects_and_no_concurrent_refresh(self): - """IdP rejects us and no sibling rotated the tokens → invalidate session.""" - from services.auth import refresh as refresh_mod - from services.auth.refresh import refresh_session_if_needed - - now = datetime.now() - stale_session = { - "id": 1, - "access_token_expires_at": now + timedelta(seconds=5), - "refresh_token_encrypted": b"old-refresh", - "last_refresh_at": None, - } - # Re-read returns the same stale row (no sibling rotation). - stale_row_from_db = dict(stale_session) - - vdb = MagicMock() - vdb.get_oidc_session_by_id_for_request = MagicMock() - vdb.get_oidc_session_by_id_for_request = AsyncMock(return_value=stale_row_from_db) - - fake_client = MagicMock() - fake_client.refresh_access_token = AsyncMock(side_effect=RuntimeError("invalid_grant")) - with ( - patch.object(refresh_mod, "get_oidc_client", return_value=fake_client), - patch.object(refresh_mod, "decrypt_token", return_value="old-refresh-plain"), - ): - out = await refresh_session_if_needed(session=stale_session, enc_key="k", auth_service=vdb) - - assert out is None - - -# --------------------------------------------------------------------------- -# is_bypass_path — regression for #359 (chainlit prefix bypass tightening) -# --------------------------------------------------------------------------- - - -def test_is_bypass_path_chainlit_subtree_only(): - """Only the actual /chainlit subtree may bypass auth — not /chainlitevil.""" - from api.middleware.auth import is_bypass_path - - # Legitimate Chainlit paths still bypass - assert is_bypass_path("/chainlit") is True - assert is_bypass_path("/chainlit/") is True - assert is_bypass_path("/chainlit/login") is True - assert is_bypass_path("/chainlit/anything/deep") is True - - # Paths that merely share the /chainlit prefix must not bypass auth - assert is_bypass_path("/chainlitevil") is False - assert is_bypass_path("/chainlit-spoof") is False - assert is_bypass_path("/chainlitX") is False diff --git a/openrag/components/auth/test_oidc_client.py b/openrag/components/auth/test_oidc_client.py deleted file mode 100644 index 1b45aa2d0..000000000 --- a/openrag/components/auth/test_oidc_client.py +++ /dev/null @@ -1,433 +0,0 @@ -"""Unit tests for oidc_client.py — uses respx to mock httpx calls.""" - -import time - -import httpx -import pytest -import pytest_asyncio -import respx -from authlib.jose import JsonWebKey -from components.auth.oidc_client import LogoutTokenClaims, OIDCClient, TokenBundle - -# --------------------------------------------------------------------------- -# Helpers — RSA test key + JWT factory -# --------------------------------------------------------------------------- - -ISSUER = "https://idp.example.com/realms/openrag" -CLIENT_ID = "openrag-client" -CLIENT_SECRET = "test-secret" -REDIRECT_URI = "https://openrag.example.com/auth/callback" -SCOPES = "openid email profile offline_access" - - -def _make_rsa_key_pair(): - """Generate an RSA-2048 key pair using authlib's JsonWebKey.""" - private = JsonWebKey.generate_key("RSA", 2048, is_private=True) - private_jwk = private.as_dict(is_private=True) - public_jwk = private.as_dict() - return private, private_jwk, public_jwk - - -# Generate once per module -_RSA_PRIVATE, _RSA_PRIVATE_JWK, _RSA_PUBLIC_JWK = _make_rsa_key_pair() -_RSA_PUBLIC_JWK["use"] = "sig" -_RSA_PUBLIC_JWK["alg"] = "RS256" -_RSA_PUBLIC_JWK["kid"] = "test-key-1" -_RSA_PRIVATE_JWK["kid"] = "test-key-1" - -JWKS_RESPONSE = {"keys": [_RSA_PUBLIC_JWK]} - -DISCOVERY_DOC = { - "issuer": ISSUER, - "authorization_endpoint": f"{ISSUER}/protocol/openid-connect/auth", - "token_endpoint": f"{ISSUER}/protocol/openid-connect/token", - "userinfo_endpoint": f"{ISSUER}/protocol/openid-connect/userinfo", - "jwks_uri": f"{ISSUER}/protocol/openid-connect/certs", - "end_session_endpoint": f"{ISSUER}/protocol/openid-connect/logout", -} - - -def _sign_jwt(payload: dict) -> str: - """Sign payload with the test RSA private key, returning a compact JWT string.""" - from authlib.jose import JsonWebToken - - header = {"alg": "RS256", "kid": "test-key-1"} - # Authlib >=1.0 requires the allowed-algorithms list on JsonWebToken. - jwt = JsonWebToken(["RS256"]) - token = jwt.encode(header, payload, _RSA_PRIVATE) - # authlib returns bytes - if isinstance(token, bytes): - return token.decode() - return token - - -def _id_token_payload(nonce: str, *, extra: dict | None = None) -> dict: - now = int(time.time()) - payload = { - "iss": ISSUER, - "sub": "user-sub-001", - "aud": CLIENT_ID, - "exp": now + 300, - "iat": now, - "nonce": nonce, - "email": "user@example.com", - } - if extra: - payload.update(extra) - return payload - - -def _logout_token_payload( - *, sub: str | None = "user-sub-001", sid: str | None = None, extra: dict | None = None -) -> dict: - now = int(time.time()) - payload = { - "iss": ISSUER, - "aud": CLIENT_ID, - "iat": now, - "jti": "logout-jti-001", - "events": {"http://schemas.openid.net/event/backchannel-logout": {}}, - } - if sub is not None: - payload["sub"] = sub - if sid is not None: - payload["sid"] = sid - if extra: - payload.update(extra) - return payload - - -# --------------------------------------------------------------------------- -# Fixture — OIDCClient with mocked httpx transport -# --------------------------------------------------------------------------- - - -@pytest_asyncio.fixture -async def client(): - """OIDCClient backed by a real httpx.AsyncClient wired to a respx MockRouter. - - respx >= 0.22 removed the top-level ``MockTransport``; use ``MockRouter`` - plus ``httpx.MockTransport(router.handler)`` instead. - """ - router = respx.MockRouter(assert_all_called=False) - http = httpx.AsyncClient(transport=httpx.MockTransport(router.handler)) - oc = OIDCClient( - issuer=ISSUER, - client_id=CLIENT_ID, - client_secret=CLIENT_SECRET, - redirect_uri=REDIRECT_URI, - scopes=SCOPES, - http_client=http, - ) - # Expose the router so individual tests can register additional routes. - oc._mock_router = router - yield oc - await oc.aclose() - - -def _setup_discovery(router: respx.MockRouter): - router.get(f"{ISSUER}/.well-known/openid-configuration").mock(return_value=httpx.Response(200, json=DISCOVERY_DOC)) - - -def _setup_jwks(router: respx.MockRouter): - router.get(f"{ISSUER}/protocol/openid-connect/certs").mock(return_value=httpx.Response(200, json=JWKS_RESPONSE)) - - -# --------------------------------------------------------------------------- -# PKCE generation tests (pure, no HTTP) -# --------------------------------------------------------------------------- - - -class TestPKCE: - def test_verifier_length(self): - verifier, _ = OIDCClient.generate_pkce_pair() - assert 43 <= len(verifier) <= 128 - - def test_challenge_is_urlsafe_base64(self): - import base64 - import hashlib - - verifier, challenge = OIDCClient.generate_pkce_pair() - expected = base64.urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()).rstrip(b"=").decode() - assert challenge == expected - - def test_unique_pairs(self): - pairs = {OIDCClient.generate_pkce_pair()[0] for _ in range(20)} - assert len(pairs) == 20 - - def test_state_and_nonce_unique(self): - states = {OIDCClient.generate_state_and_nonce()[0] for _ in range(20)} - assert len(states) == 20 - - -# --------------------------------------------------------------------------- -# Authorization URL -# --------------------------------------------------------------------------- - - -class TestBuildAuthorizationUrl: - @pytest.mark.asyncio - async def test_required_params(self, client): - _setup_discovery(client._mock_router) - url = await client.build_authorization_url(state="mystate", nonce="mynonce", code_challenge="mychallenge") - assert "response_type=code" in url - assert "client_id=openrag-client" in url - assert "state=mystate" in url - assert "nonce=mynonce" in url - assert "code_challenge=mychallenge" in url - assert "code_challenge_method=S256" in url - assert url.startswith(DISCOVERY_DOC["authorization_endpoint"]) - - -# --------------------------------------------------------------------------- -# Discovery -# --------------------------------------------------------------------------- - - -class TestDiscover: - @pytest.mark.asyncio - async def test_issuer_mismatch_raises(self, client): - bad_doc = dict(DISCOVERY_DOC, issuer="https://evil.example.com") - client._mock_router.get(f"{ISSUER}/.well-known/openid-configuration").mock( - return_value=httpx.Response(200, json=bad_doc) - ) - with pytest.raises(ValueError, match="Issuer mismatch"): - await client.discover() - - @pytest.mark.asyncio - async def test_caching(self, client): - _setup_discovery(client._mock_router) - doc1 = await client.discover() - doc2 = await client.discover() - # Same object from cache - assert doc1 is doc2 - - -# --------------------------------------------------------------------------- -# Code exchange -# --------------------------------------------------------------------------- - - -class TestExchangeCode: - @pytest.mark.asyncio - async def test_success(self, client): - _setup_discovery(client._mock_router) - _setup_jwks(client._mock_router) - - nonce = "test-nonce-abc" - id_token = _sign_jwt(_id_token_payload(nonce)) - token_response = { - "id_token": id_token, - "access_token": "at-123", - "refresh_token": "rt-456", - "expires_in": 300, - "token_type": "Bearer", - } - client._mock_router.post(f"{ISSUER}/protocol/openid-connect/token").mock( - return_value=httpx.Response(200, json=token_response) - ) - - bundle = await client.exchange_code(code="auth-code", code_verifier="verifier", expected_nonce=nonce) - assert isinstance(bundle, TokenBundle) - assert bundle.access_token == "at-123" - assert bundle.refresh_token == "rt-456" - assert bundle.claims["sub"] == "user-sub-001" - assert bundle.claims["nonce"] == nonce - - @pytest.mark.asyncio - async def test_nonce_mismatch_raises(self, client): - _setup_discovery(client._mock_router) - _setup_jwks(client._mock_router) - - id_token = _sign_jwt(_id_token_payload("correct-nonce")) - token_response = { - "id_token": id_token, - "access_token": "at", - "expires_in": 300, - "token_type": "Bearer", - } - client._mock_router.post(f"{ISSUER}/protocol/openid-connect/token").mock( - return_value=httpx.Response(200, json=token_response) - ) - - with pytest.raises(ValueError, match="nonce"): - await client.exchange_code(code="code", code_verifier="v", expected_nonce="wrong-nonce") - - # OIDC Core 1.0 §3.1.3.7: when aud lists multiple audiences, azp MUST - # be present and equal to client_id. Regression for #385. - @pytest.mark.asyncio - async def test_multi_aud_requires_matching_azp(self, client): - _setup_discovery(client._mock_router) - _setup_jwks(client._mock_router) - - nonce = "n-azp" - # Multi-aud token with the wrong (or missing) azp must be rejected - id_token = _sign_jwt( - _id_token_payload(nonce, extra={"aud": [CLIENT_ID, "other-client"], "azp": "other-client"}) - ) - token_response = { - "id_token": id_token, - "access_token": "at", - "expires_in": 300, - "token_type": "Bearer", - } - client._mock_router.post(f"{ISSUER}/protocol/openid-connect/token").mock( - return_value=httpx.Response(200, json=token_response) - ) - - with pytest.raises(ValueError, match="multi-aud"): - await client.exchange_code(code="code", code_verifier="v", expected_nonce=nonce) - - @pytest.mark.asyncio - async def test_multi_aud_with_correct_azp_passes(self, client): - _setup_discovery(client._mock_router) - _setup_jwks(client._mock_router) - - nonce = "n-azp-ok" - id_token = _sign_jwt(_id_token_payload(nonce, extra={"aud": [CLIENT_ID, "other-client"], "azp": CLIENT_ID})) - token_response = { - "id_token": id_token, - "access_token": "at", - "expires_in": 300, - "token_type": "Bearer", - } - client._mock_router.post(f"{ISSUER}/protocol/openid-connect/token").mock( - return_value=httpx.Response(200, json=token_response) - ) - - bundle = await client.exchange_code(code="code", code_verifier="v", expected_nonce=nonce) - assert bundle.claims["aud"] == [CLIENT_ID, "other-client"] - assert bundle.claims["azp"] == CLIENT_ID - - -# --------------------------------------------------------------------------- -# Token refresh -# --------------------------------------------------------------------------- - - -class TestRefreshAccessToken: - @pytest.mark.asyncio - async def test_keeps_old_refresh_token_when_omitted(self, client): - _setup_discovery(client._mock_router) - _setup_jwks(client._mock_router) - - # IdP returns no refresh_token in the response - token_response = { - "access_token": "new-at", - "expires_in": 300, - "token_type": "Bearer", - # no refresh_token - } - client._mock_router.post(f"{ISSUER}/protocol/openid-connect/token").mock( - return_value=httpx.Response(200, json=token_response) - ) - - bundle = await client.refresh_access_token("old-rt") - assert bundle.refresh_token == "old-rt" - assert bundle.access_token == "new-at" - - @pytest.mark.asyncio - async def test_uses_new_refresh_token_when_provided(self, client): - _setup_discovery(client._mock_router) - _setup_jwks(client._mock_router) - - token_response = { - "access_token": "new-at", - "refresh_token": "new-rt", - "expires_in": 300, - "token_type": "Bearer", - } - client._mock_router.post(f"{ISSUER}/protocol/openid-connect/token").mock( - return_value=httpx.Response(200, json=token_response) - ) - - bundle = await client.refresh_access_token("old-rt") - assert bundle.refresh_token == "new-rt" - - -# --------------------------------------------------------------------------- -# Userinfo -# --------------------------------------------------------------------------- - - -class TestFetchUserinfo: - @pytest.mark.asyncio - async def test_returns_userinfo(self, client): - _setup_discovery(client._mock_router) - - userinfo = {"sub": "user-sub-001", "email": "user@example.com"} - client._mock_router.get(f"{ISSUER}/protocol/openid-connect/userinfo").mock( - return_value=httpx.Response(200, json=userinfo) - ) - - result = await client.fetch_userinfo("at-123") - assert result["email"] == "user@example.com" - - -# --------------------------------------------------------------------------- -# Logout token verification -# --------------------------------------------------------------------------- - - -class TestVerifyLogoutToken: - @pytest.mark.asyncio - async def test_valid_logout_token_with_sub(self, client): - _setup_discovery(client._mock_router) - _setup_jwks(client._mock_router) - - token = _sign_jwt(_logout_token_payload(sub="user-sub-001")) - claims = await client.verify_logout_token(token) - assert isinstance(claims, LogoutTokenClaims) - assert claims.sub == "user-sub-001" - - @pytest.mark.asyncio - async def test_valid_logout_token_with_sid(self, client): - _setup_discovery(client._mock_router) - _setup_jwks(client._mock_router) - - token = _sign_jwt(_logout_token_payload(sub=None, sid="session-abc")) - claims = await client.verify_logout_token(token) - assert claims.sid == "session-abc" - assert claims.sub is None - - @pytest.mark.asyncio - async def test_missing_events_claim_raises(self, client): - _setup_discovery(client._mock_router) - _setup_jwks(client._mock_router) - - payload = _logout_token_payload() - del payload["events"] - token = _sign_jwt(payload) - with pytest.raises(ValueError, match="back-channel-logout"): - await client.verify_logout_token(token) - - @pytest.mark.asyncio - async def test_wrong_events_key_raises(self, client): - _setup_discovery(client._mock_router) - _setup_jwks(client._mock_router) - - payload = _logout_token_payload() - payload["events"] = {"http://schemas.openid.net/event/OTHER": {}} - token = _sign_jwt(payload) - with pytest.raises(ValueError, match="back-channel-logout"): - await client.verify_logout_token(token) - - @pytest.mark.asyncio - async def test_nonce_present_raises(self, client): - _setup_discovery(client._mock_router) - _setup_jwks(client._mock_router) - - payload = _logout_token_payload() - payload["nonce"] = "forbidden" - token = _sign_jwt(payload) - with pytest.raises(ValueError, match="nonce"): - await client.verify_logout_token(token) - - @pytest.mark.asyncio - async def test_missing_sub_and_sid_raises(self, client): - _setup_discovery(client._mock_router) - _setup_jwks(client._mock_router) - - token = _sign_jwt(_logout_token_payload(sub=None, sid=None)) - with pytest.raises(ValueError, match="sub or sid"): - await client.verify_logout_token(token) diff --git a/openrag/components/auth/test_session_tokens.py b/openrag/components/auth/test_session_tokens.py deleted file mode 100644 index 191688962..000000000 --- a/openrag/components/auth/test_session_tokens.py +++ /dev/null @@ -1,68 +0,0 @@ -"""Unit tests for session_tokens.py.""" - -import pytest -from components.auth.session_tokens import ( - decrypt_token, - encrypt_token, - hash_session_token, - issue_session_token, -) -from cryptography.fernet import Fernet - - -def _valid_key() -> str: - return Fernet.generate_key().decode() - - -class TestIssueSessionToken: - def test_returns_tuple_of_two_strings(self): - plain, hashed = issue_session_token() - assert isinstance(plain, str) - assert isinstance(hashed, str) - - def test_hash_is_64_hex_chars(self): - _, hashed = issue_session_token() - assert len(hashed) == 64 - assert all(c in "0123456789abcdef" for c in hashed) - - def test_tokens_are_unique(self): - tokens = {issue_session_token()[0] for _ in range(20)} - assert len(tokens) == 20 - - def test_hash_matches_plain(self): - plain, hashed = issue_session_token() - assert hash_session_token(plain) == hashed - - -class TestHashSessionToken: - def test_deterministic(self): - assert hash_session_token("abc") == hash_session_token("abc") - - def test_different_inputs_differ(self): - assert hash_session_token("abc") != hash_session_token("def") - - -class TestEncryptDecryptRoundTrip: - def test_round_trip(self): - key = _valid_key() - plaintext = "super-secret-access-token" - ciphertext = encrypt_token(plaintext, key) - assert ciphertext is not None - assert decrypt_token(ciphertext, key) == plaintext - - def test_none_plaintext_returns_none(self): - assert encrypt_token(None, _valid_key()) is None - - def test_none_ciphertext_returns_none(self): - assert decrypt_token(None, _valid_key()) is None - - def test_wrong_key_raises_value_error(self): - key1 = _valid_key() - key2 = _valid_key() - ciphertext = encrypt_token("secret", key1) - with pytest.raises(ValueError, match="decrypt"): - decrypt_token(ciphertext, key2) - - def test_invalid_key_raises_value_error(self): - with pytest.raises(ValueError, match="valid Fernet"): - encrypt_token("data", "not-a-fernet-key") diff --git a/openrag/components/auth/test_state_cookie.py b/openrag/components/auth/test_state_cookie.py deleted file mode 100644 index 4a5b8d2a2..000000000 --- a/openrag/components/auth/test_state_cookie.py +++ /dev/null @@ -1,67 +0,0 @@ -"""Unit tests for state_cookie.py.""" - -import time - -import pytest -from components.auth.state_cookie import StateCookiePayload, StateCookieSerializer - -SECRET = "test-secret-key-for-state-cookie" - - -def _serializer() -> StateCookieSerializer: - return StateCookieSerializer(SECRET) - - -def _payload() -> StateCookiePayload: - return StateCookiePayload( - state="abc123", - nonce="xyz789", - code_verifier="verifier_value", - next_url="/dashboard", - ) - - -class TestRoundTrip: - def test_dumps_loads_roundtrip(self): - ser = _serializer() - p = _payload() - token = ser.dumps(p) - result = ser.loads(token) - assert result.state == p.state - assert result.nonce == p.nonce - assert result.code_verifier == p.code_verifier - assert result.next_url == p.next_url - - def test_default_next_url(self): - ser = _serializer() - p = StateCookiePayload(state="s", nonce="n", code_verifier="v") - token = ser.dumps(p) - result = ser.loads(token) - assert result.next_url == "/" - - -class TestTampering: - def test_tampered_cookie_raises_value_error(self): - ser = _serializer() - token = ser.dumps(_payload()) - # Flip a character near the end of the token - tampered = token[:-4] + "XXXX" - with pytest.raises(ValueError, match="signature invalid"): - ser.loads(tampered) - - def test_different_secret_raises_value_error(self): - ser1 = _serializer() - ser2 = StateCookieSerializer("different-secret") - token = ser1.dumps(_payload()) - with pytest.raises(ValueError, match="signature invalid"): - ser2.loads(token) - - -class TestExpiry: - def test_expired_cookie_raises_value_error(self): - ser = _serializer() - token = ser.dumps(_payload()) - # Use max_age=0: any token older than 0 seconds is expired - time.sleep(1) - with pytest.raises(ValueError, match="expired"): - ser.loads(token, max_age=0) diff --git a/openrag/components/files.py b/openrag/components/files.py deleted file mode 100644 index a4abc2520..000000000 --- a/openrag/components/files.py +++ /dev/null @@ -1,43 +0,0 @@ -import secrets -import time -from pathlib import Path - -import aiofiles -import consts -from fastapi import UploadFile - - -def make_unique_filename(filename: str) -> Path: - ts = int(time.time() * 1000) - rand = secrets.token_hex(2) - unique_name = f"{ts}_{rand}_{filename}" - return unique_name - - -async def save_file_to_disk( - file: UploadFile, - dest_dir: Path, - chunk_size: int = consts.FILE_READ_CHUNK_SIZE, - with_random_prefix: bool = False, -) -> Path: - """ - Save file to disk by chunks, to avoid reading the whole file at once in memory. - Returns the path to the saved file. - """ - dest_dir.mkdir(parents=True, exist_ok=True) - - if with_random_prefix: - filename = make_unique_filename(file.filename) - else: - filename = file.filename - file_path = dest_dir / filename - - async with aiofiles.open(file_path, "wb") as buffer: - # Non-blocking I/O - while True: - chunk = await file.read(chunk_size) - if not chunk: - break - await buffer.write(chunk) - - return file_path diff --git a/openrag/components/indexer/__init__.py b/openrag/components/indexer/__init__.py deleted file mode 100644 index 452a565b5..000000000 --- a/openrag/components/indexer/__init__.py +++ /dev/null @@ -1,8 +0,0 @@ -"""Legacy indexer package. - -Import concrete classes from their modules directly. Keeping this package -initializer empty avoids import-time Ray/bootstrap side effects when callers -only need nested utility modules such as ``components.indexer.utils.files``. -""" - -__all__: list[str] = [] diff --git a/openrag/components/indexer/chunker/__init__.py b/openrag/components/indexer/chunker/__init__.py deleted file mode 100644 index de8090928..000000000 --- a/openrag/components/indexer/chunker/__init__.py +++ /dev/null @@ -1 +0,0 @@ -from .chunker import * diff --git a/openrag/components/indexer/chunker/chunker.py b/openrag/components/indexer/chunker/chunker.py deleted file mode 100644 index 5427118bb..000000000 --- a/openrag/components/indexer/chunker/chunker.py +++ /dev/null @@ -1,222 +0,0 @@ -"""Backward-compatibility shim — chunking primitives delegate to `openrag.core.chunking`. - -`ChunkerFactory` is config-driven; the new code uses `chunking_registry`. Both -coexist until Phase 8 cutover. - -Scheduled for removal in Phase 12. -""" - -from typing import TYPE_CHECKING, Any, ClassVar, Literal - -# Side-effect import: pre-loads the indexer-utils submodule so the legacy -# circular import between `components.utils` and `components.indexer.utils.files` -# resolves in the correct order. Removing this line breaks chunker collection. -# Slated to disappear when `components.utils` is split (Phase 6+). -from components.indexer.utils import text_sanitizer as _text_sanitizer # noqa: F401 -from components.prompts import CHUNK_CONTEXTUALIZER_PROMPT -from components.utils import detect_language, get_vlm_semaphore, load_config -from core.chunking.recursive import RecursiveSplitter as _CoreRecursiveSplitter -from core.indexing.contextualize import ChunkContextualizer as _CoreChunkContextualizer -from core.llm.llm import LLM as _CoreLLM -from core.models.chunk import Chunk as _CoreChunk -from core.models.document import ProcessedDocument, TextBlock -from core.prompts.contextualization_builder import wrap_chunk_with_context -from core.utils.logging import get_logger -from langchain_core.documents.base import Document -from langchain_core.messages import AIMessage, HumanMessage, SystemMessage - -if TYPE_CHECKING: - from components.indexer.embeddings import BaseEmbedding - from langchain_openai import ChatOpenAI - -logger = get_logger() - - -class _LangChainLLMAdapter(_CoreLLM): - """Wraps a LangChain ``ChatOpenAI`` so it satisfies the core ``LLM`` ABC.""" - - _ROLE_MAP: ClassVar[dict] = {"user": HumanMessage, "system": SystemMessage, "assistant": AIMessage} - - def __init__(self, lc_llm: "ChatOpenAI") -> None: - self._llm = lc_llm - - async def generate(self, prompt: str, **kwargs) -> str: - out = await self._llm.ainvoke(prompt) - return out.content if hasattr(out, "content") else str(out) - - async def chat(self, messages: list[dict[str, str]], **kwargs) -> str: - lc_msgs = [self._ROLE_MAP[m["role"]](content=m["content"]) for m in messages] - out = await self._llm.ainvoke(lc_msgs) - return out.content - - async def stream_chat(self, messages: list[dict[str, str]], **kwargs) -> Any: - pass # Not implemented since the contextualizer never streams. - - -def _chunks_to_documents(chunks: list, base_metadata: dict) -> list[Document]: - """Convert a list of core domain Chunks into legacy LangChain Documents. - - Reproduces the legacy metadata shape: `page`, `chunk_type`, plus the - document/partition keys the legacy code stamps onto every chunk. - """ - out: list[Document] = [] - for c in chunks: - meta = dict(c.metadata) - meta.update( - { - "file_id": c.document_id, - "partition": c.partition, - "page": c.page_number, - "chunk_type": c.chunk_type.value, - } - ) - # Preserve legacy keys that weren't lifted into core fields. - for k, v in base_metadata.items(): - meta.setdefault(k, v) - out.append(Document(page_content=c.text, metadata=meta)) - return out - - -class BaseChunker: - """Legacy chunker shell — markdown-aware splitting delegated to core. - - Subclasses configure ``self._core_splitter`` (a - `core.chunking.recursive.RecursiveSplitter`) in their `__init__`. - """ - - def __init__( - self, - chunk_size: int = 200, - chunk_overlap_rate: float = 0.2, - llm_config: dict | None = None, - contextual_retrieval: bool = False, - **kwargs, - ): - from langchain_openai import ChatOpenAI - - self.chunk_size = chunk_size - self.chunk_overlap_rate = chunk_overlap_rate - self.chunk_overlap = int(self.chunk_size * self.chunk_overlap_rate) - - self.llm = ChatOpenAI(**llm_config) - self._length_function = self.llm.get_num_tokens - - self._core_splitter: _CoreRecursiveSplitter | None = None - - self.contextual_retrieval = contextual_retrieval - if contextual_retrieval: - config = load_config() - contextualization_timeout = config.chunker.contextualization_timeout - _lc_llm = ChatOpenAI(**{**llm_config, "timeout": contextualization_timeout}) - self.contextualizer: _CoreChunkContextualizer | None = _CoreChunkContextualizer( - llm=_LangChainLLMAdapter(_lc_llm), - system_prompt=CHUNK_CONTEXTUALIZER_PROMPT, - timeout_seconds=contextualization_timeout, - max_concurrent=config.chunker.max_concurrent_contextualization, - semaphore=get_vlm_semaphore(), - ) - else: - self.contextualizer = None - - async def _apply_contextualization( - self, - chunks: list[Document], - lang: Literal["en", "fr"] = "en", - filename: str = "", - ) -> list[Document]: - """Apply contextualization if enabled.""" - if not self.contextual_retrieval or len(chunks) < 2: - return [ - Document( - page_content=wrap_chunk_with_context(c.page_content, filename), - metadata=c.metadata, - ) - for c in chunks - ] - - core_chunks = [_CoreChunk.from_langchain(c) for c in chunks] - contextualized = await self.contextualizer.contextualize(core_chunks, filename=filename, lang=lang) - return [c.to_langchain(with_id=False) for c in contextualized] - - def _get_chunks(self, content: str, metadata: dict | None = None, log=None) -> list[Document]: - log = log or logger - metadata = metadata or {} - partition = metadata.get("partition", "default") - - doc = ProcessedDocument( - document_id=metadata.get("file_id", ""), - text_blocks=[TextBlock(text=content)], - metadata=metadata, - ) - chunks = self._core_splitter.chunk(doc, partition=partition) - if not chunks: - log.warning("No chunks created. Content is empty or image is not informative.") - return [] - return _chunks_to_documents(chunks, base_metadata=metadata) - - async def split_document(self, doc: Document, task_id: str | None = None) -> list[Document]: - """Split document into chunks with optional contextualization.""" - metadata = doc.metadata - filename = metadata.get("filename", "") - log = logger.bind( - file_id=metadata.get("file_id"), - partition=metadata.get("partition"), - task_id=task_id, - ) - log.info("Starting document chunking") - - detected_lang = detect_language(text=doc.page_content) - - chunks = self._get_chunks(doc.page_content.strip(), metadata, log=log) - - if chunks: - log.info( - "Contextualizing chunks", - apply_contextualization=self.contextual_retrieval, - ) - chunks = await self._apply_contextualization(chunks, lang=detected_lang, filename=filename) - log.info("Document chunking completed") - return chunks - else: - return [] - - -class RecursiveSplitter(BaseChunker): - def __init__( - self, - chunk_size=200, - chunk_overlap_rate=0.2, - llm_config=None, - contextual_retrieval=False, - **kwargs, - ): - super().__init__(chunk_size, chunk_overlap_rate, llm_config, contextual_retrieval, **kwargs) - self._core_splitter = _CoreRecursiveSplitter( - chunk_size=self.chunk_size, - chunk_overlap_rate=self.chunk_overlap_rate, - length_function=self._length_function, - ) - - -class ChunkerFactory: - CHUNKERS = { - "recursive_splitter": RecursiveSplitter, - } - - @staticmethod - def create_chunker( - config, - embedder: "BaseEmbedding | None" = None, - ) -> BaseChunker: - chunker_params = config.chunker.model_dump() - name = chunker_params.pop("name") - - chunker_cls: BaseChunker = ChunkerFactory.CHUNKERS.get(name) - - if not chunker_cls: - raise ValueError( - f"Chunker '{name}' is not recognized. Available chunkers: {list(ChunkerFactory.CHUNKERS.keys())}" - ) - - chunker_params["llm_config"] = config.vlm.model_dump() - return chunker_cls(**chunker_params) diff --git a/openrag/components/indexer/chunker/test_chunking.py b/openrag/components/indexer/chunker/test_chunking.py deleted file mode 100644 index c3ffd9bd2..000000000 --- a/openrag/components/indexer/chunker/test_chunking.py +++ /dev/null @@ -1,232 +0,0 @@ -# from .utils import split_md_elements - -from components.indexer.chunker.utils import ( - MDElement, - chunk_table, - get_chunk_page_number, - split_md_elements, -) -from components.indexer.utils.text_sanitizer import clean_markdown_table_spacing - - -class TestSplitMdElements: - """Test suite for split_md_elements function.""" - - def test_simple_text_only(self): - """Test parsing markdown with only text content.""" - md_text = "This is a simple paragraph.\n\nAnother paragraph here." - elements = split_md_elements(md_text) - - assert len(elements) == 1 - assert elements[0].type == "text" - assert md_text == elements[0].content - - def test_single_table(self): - """Test parsing a single markdown table.""" - - md_text = "Some text before.\n\n| Header 1 | Header 2 |\n|----------|----------|\n| Cell 1 | Cell 2 |\n| Cell 3 | Cell 4 |\n\nSome text after." - elements = split_md_elements(md_text) - - # Should have: text, table, text - assert len(elements) == 3 - assert elements[0].type == "text" - assert elements[1].type == "table" - assert elements[2].type == "text" - assert "Header 1" in elements[1].content - - def test_single_image(self): - """Test parsing a single image description.""" - md_text = """ -Text before image. - - -A beautiful sunset over the ocean. - - -Text after image.""" - elements = split_md_elements(md_text) - - assert len(elements) == 3 - assert elements[0].type == "text" - assert elements[1].type == "image" - assert elements[2].type == "text" - assert "sunset" in elements[1].content - - def test_table_inside_image_description(self): - """Test that tables inside image descriptions are ignored.""" - md_text = """ - -This image contains a table: -| Col 1 | Col 2 | -|-------|-------| -| A | B | - - -Outside table: -| Real 1 | Real 2 | -|--------|--------| -| X | Y | -""" - elements = split_md_elements(md_text) - - # Should have: image, text, table - assert len(elements) == 3 - assert elements[0].type == "image" - assert elements[1].type == "text" - assert elements[2].type == "table" - # The table inside image should not be parsed separately - table_elements = [e for e in elements if e.type == "table"] - assert len(table_elements) == 1 - assert "Real 1" in table_elements[0].content - - def test_page_markers_with_table(self): - """Test page number assignment for tables.""" - md_text = """text on page 1. -[PAGE_1] -Text on page 2. - -| Header 1 | Header 2 | -|----------|----------| -| Data 1 | Data 2 | - -[PAGE_2] -More content. -""" - elements = split_md_elements(md_text) - assert len(elements) == 3 # text, table, text - - table_elements = [e for e in elements if e.type == "table"] - assert len(table_elements) == 1 - assert table_elements[0].page_number == 2 - - def test_page_markers_with_images(self): - """Test page number assignment for images.""" - md_text = """ -[PAGE_1] -[PAGE_2] - -Image on page 3. - -""" - elements = split_md_elements(md_text) - - image_elements = [e for e in elements if e.type == "image"] - assert len(image_elements) == 1 - assert image_elements[0].page_number == 3 - - -class TestGetChunkPageNumber: - """Test suite for get_chunk_page_number function.""" - - def test_no_page_markers(self): - """Test chunk with no page markers.""" - chunk = "Just some plain text content." - result = get_chunk_page_number(chunk, previous_chunk_ending_page=1) - - assert result["start_page"] == 1 - assert result["end_page"] == 1 - - def test_starts_with_marker(self): - """Test chunk starting with a page marker.""" - chunk = "[PAGE_2]Content on page 3." - result = get_chunk_page_number(chunk, previous_chunk_ending_page=1) - - assert result["start_page"] == 3 - assert result["end_page"] == 3 - - def test_ends_with_marker(self): - """Test chunk ending with a page marker.""" - chunk = "Content on page 1.[PAGE_1]" - result = get_chunk_page_number(chunk, previous_chunk_ending_page=1) - - assert result["start_page"] == 1 - assert result["end_page"] == 1 - - def test_marker_in_middle(self): - """Test chunk with marker in the middle.""" - chunk = "Start on page 1.[PAGE_1]End on page 2." - result = get_chunk_page_number(chunk, previous_chunk_ending_page=1) - - assert result["start_page"] == 1 - assert result["end_page"] == 2 - - -class TestCleanMarkdownTableSpacing: - """Test suite for clean_markdown_table_spacing function.""" - - def test_extra_spaces_in_cells(self): - """Test trimming excessive spaces within cells.""" - table = "| Header 1 | Header 2 |\n|-------------|-------------|\n| Cell 1 | Cell 2 |" - result = clean_markdown_table_spacing(table) - - assert result == "| Header 1 | Header 2 |\n| ------------- | ------------- |\n| Cell 1 | Cell 2 |" - - def test_inconsistent_spacing(self): - """Test normalizing inconsistent spacing across rows.""" - table = "|Header1|Header2|\n|---|---|\n| A |B|" - result = clean_markdown_table_spacing(table) - - assert result == "| Header1 | Header2 |\n| --- | --- |\n| A | B |" - - def test_empty_cells(self): - """Test handling of empty cells.""" - table = "| Col1 | Col2 | Col3 |\n|------|------|------|\n| Data | | More |\n| | Data | |" - result = clean_markdown_table_spacing(table) - - assert result == "| Col1 | Col2 | Col3 |\n| ------ | ------ | ------ |\n| Data | | More |\n| | Data | |" - - def test_multiline_spacing(self): - """Test table with varying amounts of whitespace.""" - table = "| A | B | C |\n|------|--------|----------|\n|1|2|3|" - result = clean_markdown_table_spacing(table) - - assert result == "| A | B | C |\n| ------ | -------- | ---------- |\n| 1 | 2 | 3 |" - - -class TestChunkTable: - """Test suite for chunk_table function.""" - - def mock_length_function(self, text): - """Mock function that estimates token count (~4 chars per token).""" - return len(text) // 4 - - def test_small_table_no_chunking(self): - """Test that a small table remains as a single chunk.""" - table_content = "| Name | Age |\n|------|-----|\n| John | 30 |\n| Jane | 25 |" - table_element = MDElement(type="table", content=table_content, page_number=1) - - chunks = chunk_table(table_element, chunk_size=1000, length_function=self.mock_length_function) - - assert len(chunks) == 1 - assert chunks[0].type == "table" - assert chunks[0].page_number == 1 - assert "John" in chunks[0].content - assert "Jane" in chunks[0].content - - def test_chunking_preserves_groups(self): - """Test that country groups are not split mid-group.""" - header = "| Country | Strategy | Goals |" - g1 = "| USA | Cyber | Goal 1 |\n| | | Goal 2 |\n| | | Goal 3 |" - - g2 = "| Mexico | Defense | Goal X |\n| | | Goal Y |\n| | | Goal Z |" - - table = f"{header}\n|----|----|----|\n{g1}\n{g2}\n" - table_element = MDElement(type="table", content=table, page_number=2) - - table_token_length = self.mock_length_function(table) - chunk_size = table_token_length // 2 # to enforce chunking into 2 chunks - - # Force chunking with small chunk size - chunks = chunk_table( - table_element, - chunk_size=chunk_size, - length_function=self.mock_length_function, - ) - - assert len(chunks) == 2 # Should be chunked - assert all(chunk.type == "table" for chunk in chunks) - assert all(header in chunk.content for chunk in chunks) # All table chunks should have the header part - - # check that groups are intact - assert "USA" in chunks[0].content - assert "Mexico" in chunks[1].content diff --git a/openrag/components/indexer/chunker/utils.py b/openrag/components/indexer/chunker/utils.py deleted file mode 100644 index b07753233..000000000 --- a/openrag/components/indexer/chunker/utils.py +++ /dev/null @@ -1,33 +0,0 @@ -"""Backward-compatibility shim — re-exports from `openrag.core.chunking.markdown_utils`. - -The implementation moved to `openrag/core/chunking/markdown_utils.py` in -Phase 5B. New code should import from there directly. This file is kept -so existing legacy imports keep working until the consumers migrate; -scheduled for removal in Phase 12. -""" - -from core.chunking.markdown_utils import ( - IMAGE_RE, - PAGE_RE, - TABLE_RE, - MDElement, - chunk_table, - get_chunk_page_number, - get_page_number, - parse_markdown_table, - span_inside, - split_md_elements, -) - -__all__ = [ - "IMAGE_RE", - "MDElement", - "PAGE_RE", - "TABLE_RE", - "chunk_table", - "get_chunk_page_number", - "get_page_number", - "parse_markdown_table", - "span_inside", - "split_md_elements", -] diff --git a/openrag/components/indexer/embeddings/__init__.py b/openrag/components/indexer/embeddings/__init__.py deleted file mode 100644 index a1d35852c..000000000 --- a/openrag/components/indexer/embeddings/__init__.py +++ /dev/null @@ -1,23 +0,0 @@ -from services.inference.vllm_client import VLLMEmbedder # noqa: F401 - -from .base import BaseEmbedding -from .openai import _ShimOpenAIEmbedding as OpenAIEmbedding - -EMBEDDER_MAPPING = { - "openai": OpenAIEmbedding, -} - - -class EmbeddingFactory: - @staticmethod - def get_embedder(embeddings_config) -> BaseEmbedding: - provider = embeddings_config.provider - embedder_class = EMBEDDER_MAPPING.get(provider, None) - - if not embedder_class: - raise ValueError(f"Unsupported embedding provider: {provider}") - - return embedder_class(embeddings_config) - - -__all__ = ["BaseEmbedding", "EmbeddingFactory", "OpenAIEmbedding"] diff --git a/openrag/components/indexer/embeddings/base.py b/openrag/components/indexer/embeddings/base.py deleted file mode 100644 index dd606782b..000000000 --- a/openrag/components/indexer/embeddings/base.py +++ /dev/null @@ -1,20 +0,0 @@ -from langchain_core.documents.base import Document -from langchain_core.embeddings import Embeddings - - -class BaseEmbedding(Embeddings): - """Base class for all embedding models.""" - - @property - def embedding_dimension(self) -> int: - """ - Returns the dimension of the embedding vector. - This is used to validate the vector size during indexing and searching. - """ - raise NotImplementedError("This method should be implemented by subclasses.") - - def embed_documents(self, texts: list[str | Document]) -> list[list[float]]: - return super().embed_documents(texts) - - def embed_query(self, text: str) -> list[float]: - return super().embed_query(text) diff --git a/openrag/components/indexer/embeddings/openai.py b/openrag/components/indexer/embeddings/openai.py deleted file mode 100644 index 1083b78c4..000000000 --- a/openrag/components/indexer/embeddings/openai.py +++ /dev/null @@ -1,155 +0,0 @@ -"""Backward-compatibility shim — delegates to services.inference.vllm_client. - -All new code should import directly from ``services.inference.vllm_client``. -""" - -import asyncio -from concurrent.futures import ThreadPoolExecutor - -import httpx -import openai -from core.config.endpoints import EmbedderConfig -from core.utils.exceptions import ( - EmbeddingAPIError, - EmbeddingResponseError, - UnexpectedEmbeddingError, -) -from core.utils.logging import get_logger -from langchain_core.documents.base import Document -from openai import OpenAI -from services.inference.vllm_client import VLLMEmbedder # noqa: F401 - -from .base import BaseEmbedding - -logger = get_logger() - - -_SYNC_POOL = ThreadPoolExecutor(max_workers=1) - - -def _run_sync(coro): - """Run an async coroutine from sync code, safe inside a running event loop (e.g. Ray).""" - return _SYNC_POOL.submit(asyncio.run, coro).result() - - -def _normalize_texts(texts: list[str | Document]) -> list[str]: - return [item.page_content if isinstance(item, Document) else item for item in texts] - - -class _ShimOpenAIEmbedding(BaseEmbedding): - """Legacy shim — delegates to ``VLLMEmbedder`` for actual HTTP transport. - - Preserves the sync ``embed_documents``/``embed_query`` contract expected by - ``vectordb.py`` (via LangChain's ``aembed_documents`` thread wrapper) while - using VLLMEmbedder's long-lived async httpx pool under the hood. - """ - - def __init__(self, embeddings_config: EmbedderConfig): - self._delegate = VLLMEmbedder( - endpoint=embeddings_config.base_url, - model_name=embeddings_config.model_name, - max_model_len=embeddings_config.max_model_len, - api_key=embeddings_config.api_key, - ) - - @property - def embedding_dimension(self) -> int: - # Probe once if unknown — legacy callers (e.g. MilvusDB schema creation) read - # this before any embed() call, but VLLMEmbedder only learns its dimension - # from a real response. The probe must run on a one-off sync httpx.Client: - # asyncio.run() here would tear down the loop and leave the delegate's - # long-lived AsyncClient pool with stale connections, breaking the next - # real async call with "Event loop is closed". - try: - return self._delegate.dimension - except RuntimeError: - pass - body: dict = {"model": self._delegate._model, "input": ["dim-probe"]} - if self._delegate._max_model_len is not None: - body["truncate_prompt_tokens"] = self._delegate._max_model_len - with httpx.Client(timeout=30.0, headers=dict(self._delegate._client.headers)) as client: - resp = client.post(f"{self._delegate._endpoint}/embeddings", json=body) - resp.raise_for_status() - self._delegate._dimension = len(resp.json()["data"][0]["embedding"]) - return self._delegate._dimension - - def embed_documents(self, texts: list[str | Document]) -> list[list[float]]: - if not texts: - return [] - return _run_sync(self._delegate.embed(_normalize_texts(texts))) - - async def aembed_documents(self, texts: list[str | Document]) -> list[list[float]]: - if not texts: - return [] - return await self._delegate.embed(_normalize_texts(texts)) - - def embed_query(self, text: str) -> list[float]: - return _run_sync(self._delegate.embed_single(text)) - - async def aembed_query(self, text: str) -> list[float]: - return await self._delegate.embed_single(text) - - -class OpenAIEmbedding(BaseEmbedding): - """Legacy OpenAI embedding wrapper. New code should use VLLMEmbedder (via DI).""" - - def __init__(self, embeddings_config): - self.embedding_model = embeddings_config.model_name - self.base_url = embeddings_config.base_url - self.api_key = embeddings_config.api_key - self.max_model_len = embeddings_config.max_model_len - self._sync_client = OpenAI(base_url=self.base_url, api_key=self.api_key) - - @property - def embedding_dimension(self) -> int: - try: - output = self.embed_documents([Document(page_content="test")]) - return len(output[0]) - except Exception: - raise - - def embed_documents(self, texts: list[str | Document]) -> list[list[float]]: - if isinstance(texts[0], Document): - texts = [doc.page_content for doc in texts] - - try: - response = self._sync_client.embeddings.create( - model=self.embedding_model, - input=texts, - extra_body={"truncate_prompt_tokens": self.max_model_len}, - ) - return [vector.embedding for vector in response.data] - - except openai.APIError as e: - logger.error("API error in embed_documents", error=str(e)) - raise EmbeddingAPIError( - f"OpenAI API error during document embedding: {e!s}", - model_name=self.embedding_model, - base_url=self.base_url, - error=str(e), - ) - - except (IndexError, AttributeError) as e: - logger.error("Error while accessing embedding data", error=str(e)) - raise EmbeddingResponseError( - "Failed to retrieve document embeddings due to unexpected response format.", - model_name=self.embedding_model, - base_url=self.base_url, - error=str(e), - ) - - except Exception as e: - logger.exception("Unexpected error while embedding documents", error=str(e)) - raise UnexpectedEmbeddingError( - f"Failed to embed documents: {e!s}", - model_name=self.embedding_model, - base_url=self.base_url, - error=str(e), - ) - - def embed_query(self, text: str) -> list[float]: - try: - output = self.embed_documents([Document(page_content=text)]) - return output[0] - except Exception: - raise diff --git a/openrag/components/indexer/utils/__init__.py b/openrag/components/indexer/utils/__init__.py deleted file mode 100644 index e92207d7e..000000000 --- a/openrag/components/indexer/utils/__init__.py +++ /dev/null @@ -1,2 +0,0 @@ -from .files import * -from .text_sanitizer import * diff --git a/openrag/components/indexer/utils/files.py b/openrag/components/indexer/utils/files.py deleted file mode 100644 index 73b407fc3..000000000 --- a/openrag/components/indexer/utils/files.py +++ /dev/null @@ -1,89 +0,0 @@ -import re -import secrets -import time -from datetime import UTC, datetime -from pathlib import Path - -import aiofiles -import consts -from fastapi import HTTPException, UploadFile, status - - -def sanitize_filename(filename: str) -> str: - # Split filename into name and extension - path = Path(filename) - name = path.stem - ext = path.suffix - - # Remove special characters (keep only word characters and hyphens temporarily) - name = re.sub(r"[^\w\-]", "_", name) - - # Replace hyphens with underscores - name = name.replace("-", "_") - - # Collapse multiple underscores - name = re.sub(r"_+", "_", name) - - # Remove leading/trailing underscores - name = name.strip("_") - - # Reconstruct filename - return name + ext - - -def make_unique_filename(filename: str) -> Path: - ts = int(time.time() * 1000) - rand = secrets.token_hex(2) - unique_name = f"{ts}_{rand}_{filename}" - return unique_name - - -async def save_file_to_disk( - file: UploadFile, - dest_dir: Path, - chunk_size: int = consts.FILE_READ_CHUNK_SIZE, - with_random_prefix: bool = False, -) -> Path: - """ - Save file to disk by chunks, to avoid reading the whole file at once in memory. - Returns the path to the saved file. - """ - dest_dir.mkdir(parents=True, exist_ok=True) - - if with_random_prefix: - filename = make_unique_filename(file.filename) - else: - filename = file.filename - file_path = dest_dir / filename - - async with aiofiles.open(file_path, "wb") as buffer: - # Non-blocking I/O - while True: - chunk = await file.read(chunk_size) - if not chunk: - break - await buffer.write(chunk) - - return file_path - - -def extract_temporal_fields(metadata: dict, temporal_fields: list) -> dict: - result = {} - for field in temporal_fields: - if field not in metadata or metadata[field] is None: - continue - - datetime_str = metadata[field] - try: - # Try parsing the provided datetime to ensure it's valid - d = datetime.fromisoformat(datetime_str) - if d.tzinfo is None: - d = d.replace(tzinfo=UTC) - result[field] = d.isoformat() - except Exception: - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail=f"Invalid ISO 8601 datetime field ({datetime_str}) for field '{field}'.", - ) - - return result diff --git a/openrag/components/indexer/utils/test_text_sanitizer.py b/openrag/components/indexer/utils/test_text_sanitizer.py deleted file mode 100644 index d26f2be22..000000000 --- a/openrag/components/indexer/utils/test_text_sanitizer.py +++ /dev/null @@ -1,200 +0,0 @@ -""" -Tests for text sanitization utilities. -""" - -from .text_sanitizer import ( - clean_markdown_table_spacing, - sanitize_extracted_text, - sanitize_text, -) - - -class TestSanitizeText: - """Test suite for sanitize_text function.""" - - def test_basic_text_unchanged(self): - """Test that normal text passes through unchanged.""" - text = "Hello world. This is a test." - result = sanitize_text(text) - assert result == text - - def test_remove_excessive_spaces(self): - """Test removal of excessive spaces.""" - text = "Hello world with many spaces" - result = sanitize_text(text) - assert result == "Hello world with many spaces" - - def test_remove_tabs(self): - """Test conversion of tabs to single space.""" - text = "Hello\t\tworld\twith\t\t\ttabs" - result = sanitize_text(text) - assert result == "Hello world with tabs" - - def test_limit_consecutive_newlines(self): - """Test limiting consecutive newlines.""" - text = "Line 1\n\n\n\n\nLine 2" - result = sanitize_text(text, max_consecutive_newlines=2) - assert result == "Line 1\n\nLine 2" - - def test_remove_control_characters(self): - """Test removal of control characters.""" - text = "Hello\x00\x01\x02world" - result = sanitize_text(text) - assert result == "Helloworld" - - def test_remove_zero_width_characters(self): - """Test removal of zero-width characters.""" - text = "Hello\u200b\u200c\u200dworld" - result = sanitize_text(text) - assert result == "Helloworld" - - def test_preserve_newlines_and_tabs_when_configured(self): - """Test that newlines are preserved.""" - text = "Line 1\nLine 2\nLine 3" - result = sanitize_text(text, normalize_whitespace=False) - assert "\n" in result - - def test_trim_leading_trailing_whitespace(self): - """Test removal of leading and trailing whitespace.""" - text = " Hello world \n\n" - result = sanitize_text(text) - assert result == "Hello world" - - def test_normalize_line_breaks(self): - """Test normalization of different line break styles.""" - text = "Line 1\r\nLine 2\rLine 3\nLine 4" - result = sanitize_text(text) - assert result == "Line 1\nLine 2\nLine 3\nLine 4" - - def test_remove_leading_spaces_from_lines(self): - """Test removal of leading spaces from each line.""" - text = "Line 1\n Line 2 with spaces\n Line 3" - result = sanitize_text(text) - assert result == "Line 1\nLine 2 with spaces\nLine 3" - - def test_remove_trailing_spaces_from_lines(self): - """Test removal of trailing spaces from each line.""" - text = "Line 1 \nLine 2 with spaces \nLine 3 " - result = sanitize_text(text) - assert result == "Line 1\nLine 2 with spaces\nLine 3" - - def test_empty_string(self): - """Test handling of empty string.""" - result = sanitize_text("") - assert result == "" - - def test_complex_mixed_issues(self): - """Test handling of multiple issues simultaneously.""" - text = " Hello world\x00\t\twith\n\n\n\nmany\u200bissues \n" - result = sanitize_text(text) - assert result == "Hello world with\n\nmanyissues" - - def test_unicode_normalization(self): - """Test unicode normalization.""" - # Decomposed form: é as e + combining acute accent - text_decomposed = "café\u0301" # cafe with combining accent - result = sanitize_text(text_decomposed, normalize_unicode=True) - # Should normalize to composed form - assert "é" in result or result == "café" - - def test_disable_whitespace_normalization(self): - """Test that whitespace normalization can be disabled.""" - text = "Hello world" - result = sanitize_text(text, normalize_whitespace=False) - assert " " in result - - def test_disable_control_char_removal(self): - """Test that control character removal can be disabled.""" - text = "Hello\x00world" - result = sanitize_text(text, remove_control_chars=False) - assert "\x00" in result - - def test_no_max_consecutive_newlines(self): - """Test unlimited consecutive newlines.""" - text = "Line 1\n\n\n\n\nLine 2" - result = sanitize_text(text, max_consecutive_newlines=0) - assert result.count("\n") == 5 - - -class TestCleanMarkdownTableSpacing: - """Test suite for clean_markdown_table_spacing function.""" - - def test_extra_spaces_in_cells(self): - """Test trimming excessive spaces within cells.""" - table = "| Header 1 | Header 2 |\n|-------------|-------------|\n| Cell 1 | Cell 2 |" - result = clean_markdown_table_spacing(table) - - assert result == "| Header 1 | Header 2 |\n| ------------- | ------------- |\n| Cell 1 | Cell 2 |" - - def test_inconsistent_spacing(self): - """Test normalizing inconsistent spacing across rows.""" - table = "|Header1|Header2|\n|---|---|\n| A |B|" - result = clean_markdown_table_spacing(table) - - assert result == "| Header1 | Header2 |\n| --- | --- |\n| A | B |" - - def test_empty_cells(self): - """Test handling of empty cells.""" - table = "| Col1 | Col2 | Col3 |\n|------|------|------|\n| Data | | More |\n| | Data | |" - result = clean_markdown_table_spacing(table) - - assert result == "| Col1 | Col2 | Col3 |\n| ------ | ------ | ------ |\n| Data | | More |\n| | Data | |" - - def test_multiline_spacing(self): - """Test table with varying amounts of whitespace.""" - table = "| A | B | C |\n|------|--------|----------|\n|1|2|3|" - result = clean_markdown_table_spacing(table) - - assert result == "| A | B | C |\n| ------ | -------- | ---------- |\n| 1 | 2 | 3 |" - - -class TestSanitizeExtractedText: - """Test suite for sanitize_extracted_text convenience function.""" - - def test_applies_all_default_sanitizations(self): - """Test that all default sanitizations are applied.""" - text = " Hello world\x00\t\twith\n\n\n\nmany\u200bissues \n" - result = sanitize_extracted_text(text) - - # Should have normalized spaces - assert " " not in result - # Should have removed control chars - assert "\x00" not in result - # Should have removed zero-width chars - assert "\u200b" not in result - # Should have limited newlines to 2 - assert "\n\n\n" not in result - # Should be trimmed - assert not result.startswith(" ") - assert not result.endswith(" ") - - def test_basic_extraction_scenario(self): - """Test a realistic extraction scenario.""" - # Simulating text extracted from a PDF with various artifacts - text = """ - Document Title - - - - This is a paragraph with excessive spaces and some\t\ttabs. - - Another paragraph here. - - - Some \x00control\x01 characters\x02 too. - - And zero-width\u200bspaces. - """ # noqa: W293 - - result = sanitize_extracted_text(text) - - # Check that text is cleaned properly - assert " " not in result # No excessive spaces - assert "\t\t" not in result # No multiple tabs - assert "\x00" not in result # No control chars - assert "\u200b" not in result # No zero-width spaces - # Should have at most 2 consecutive newlines - assert "\n\n\n" not in result - # Should still have content structure - assert "Document Title" in result - assert "paragraph" in result diff --git a/openrag/components/indexer/utils/text_sanitizer.py b/openrag/components/indexer/utils/text_sanitizer.py deleted file mode 100644 index ef14f6f8d..000000000 --- a/openrag/components/indexer/utils/text_sanitizer.py +++ /dev/null @@ -1,7 +0,0 @@ -# Re-export from canonical location for backward compatibility. -# New code should import from `core.utils.text` directly. -from core.utils.text import ( # noqa: F401 - clean_markdown_table_spacing, - sanitize_extracted_text, - sanitize_text, -) diff --git a/openrag/components/llm.py b/openrag/components/llm.py deleted file mode 100644 index a2873941c..000000000 --- a/openrag/components/llm.py +++ /dev/null @@ -1,149 +0,0 @@ -"""Backward-compatibility shim — delegates to services.inference.vllm_client. - -All new code should import directly from ``services.inference.vllm_client``. -""" - -import copy -import json -import warnings - -import httpx -from config.models import LLMConfig -from core.utils.logging import get_logger -from services.inference.vllm_client import VLLMClient # noqa: F401 - -logger = get_logger() - - -class _LLMShim: - """Legacy shim — delegates to ``VLLMClient`` for retry, circuit breaker, - and connection pooling while preserving the generator-based interface.""" - - def __init__(self, llm_config: LLMConfig, logger=None): - warnings.warn( - "components.llm.LLM is deprecated — use services.inference.vllm_client.VLLMClient", - DeprecationWarning, - stacklevel=2, - ) - self.logger = logger - config_kwargs = {k: v for k, v in llm_config.model_dump().items() if k not in ("api_key", "base_url", "model")} - self._delegate = VLLMClient( - endpoint=llm_config.base_url, - model_name=llm_config.model, - api_key=llm_config.api_key, - **config_kwargs, - ) - - async def completions(self, request: dict): - prompt = request.pop("prompt") - response = await self._delegate.generate(prompt, **request) - yield response - - async def chat_completion(self, request: dict): - messages = request.pop("messages") - stream = request.pop("stream", False) - - if stream: - async for line in self._delegate.stream_chat(messages, **request): - yield line - else: - resp_dict = await self._delegate.chat(messages, **request) - yield resp_dict - - -class LLM: - """Legacy LLM wrapper. New code should use VLLMClient (via DI) instead.""" - - def __init__(self, llm_config, logger=None): - self.logger = logger - default_llm_config = llm_config.model_dump() - self._api_key = default_llm_config.pop("api_key", None) - self._base_url = default_llm_config.pop("base_url", None) - self.default_llm_config = default_llm_config - - self.headers = { - "Content-Type": "application/json", - "Authorization": f"Bearer {self._api_key}", - } - - def _extract_llm_overrides(self, request: dict): - metadata = request.get("metadata") or {} - llm_override = metadata.pop("llm_override", None) or {} - - request.pop("model") - payload = copy.deepcopy(self.default_llm_config) - payload.update(request) - - if llm_override.get("model"): - payload["model"] = llm_override["model"] - - base_url = (llm_override.get("base_url") or self._base_url).rstrip("/") - api_key = llm_override.get("api_key") or self._api_key - headers = { - "Content-Type": "application/json", - "Authorization": f"Bearer {api_key}", - } - - return payload, base_url, headers - - async def completions(self, request: dict): - payload, base_url, headers = self._extract_llm_overrides(request) - - timeout = httpx.Timeout(4 * 10) - async with httpx.AsyncClient(timeout=timeout) as client: - try: - response = await client.post( - url=f"{base_url}/completions", - headers=headers, - json=payload, - ) - response.raise_for_status() - data = response.json() - yield data - except httpx.HTTPStatusError as e: - error_detail = e.response.text - raise ValueError(f"LLM API error ({e.response.status_code}): {error_detail}") - except json.JSONDecodeError as e: - raise ValueError(f"Invalid JSON in API response: {str(e)}") - - async def chat_completion(self, request: dict): - payload, base_url, headers = self._extract_llm_overrides(request) - stream = payload["stream"] - - timeout = httpx.Timeout(4 * 60) - async with httpx.AsyncClient(timeout=timeout) as client: - if stream: - try: - async with client.stream( - "POST", - url=f"{base_url}/chat/completions", - headers=headers, - json=payload, - ) as response: - if response.status_code >= 400: - await response.aread() - error_detail = response.text - raise ValueError(f"LLM API error ({response.status_code}): {error_detail}") - async for line in response.aiter_lines(): - yield line - except ValueError: - raise - except Exception as e: - logger.error(f"Error while streaming chat completion: {str(e)}") - raise - - else: - try: - response = await client.post( - url=f"{base_url}/chat/completions", - headers=headers, - json=payload, - ) - response.raise_for_status() - data = response.json() - yield data - except httpx.HTTPStatusError as e: - error_detail = e.response.text - raise ValueError(f"LLM API error ({e.response.status_code}): {error_detail}") - except json.JSONDecodeError as e: - raise ValueError(f"Invalid JSON in API response: {str(e)}") diff --git a/openrag/components/prompts/__init__.py b/openrag/components/prompts/__init__.py deleted file mode 100644 index 25d4532cb..000000000 --- a/openrag/components/prompts/__init__.py +++ /dev/null @@ -1 +0,0 @@ -from .prompts import * diff --git a/openrag/components/prompts/prompts.py b/openrag/components/prompts/prompts.py deleted file mode 100644 index 45c5e513b..000000000 --- a/openrag/components/prompts/prompts.py +++ /dev/null @@ -1,42 +0,0 @@ -"""Backward-compatibility shim — delegates to `openrag.core.prompts.template_loader`. - -The disk-based template loader moved to -`openrag/core/prompts/template_loader.py` in Phase 5C. This module is -retained for legacy imports of `load_prompt(...)` and the eagerly-loaded -SYS_PROMPT_TMPLT / *_PROMPT constants until consumers migrate; -scheduled for removal in Phase 12. - -The new function takes (prompts_dir, mapping, key) explicitly; this -shim's `load_prompt(key)` resolves the first two from the cached -config, matching the legacy call site shape. -""" - -from pathlib import Path - -from config import load_config -from core.prompts.template_loader import load_template_by_key - - -def load_prompt( - prompt_name: str, - prompts_dir: Path | None = None, - prompt_mapping=None, -) -> str: - config = load_config() - prompts_dir = prompts_dir or config.paths.prompts_dir - prompt_mapping = prompt_mapping or config.prompts - return load_template_by_key(prompts_dir, prompt_mapping, prompt_name) - - -# Eagerly-loaded prompt strings — preserved for legacy callers that -# import these names directly. New code should call `load_template_by_key` -# (or `load_template`) on demand instead. -SYS_PROMPT_TMPLT = load_prompt("sys_prompt") -QUERY_CONTEXTUALIZER_PROMPT = load_prompt("query_contextualizer") -CHUNK_CONTEXTUALIZER_PROMPT = load_prompt("chunk_contextualizer") -IMAGE_DESCRIBER = load_prompt("image_describer") - -HYDE_PROMPT = load_prompt("hyde") -MULTI_QUERY_PROMPT = load_prompt("multi_query") - -SPOKEN_STYLE_ANSWER_PROMPT = load_prompt("spoken_style_answer") diff --git a/openrag/components/ray_utils.py b/openrag/components/ray_utils.py deleted file mode 100644 index 83ec534b1..000000000 --- a/openrag/components/ray_utils.py +++ /dev/null @@ -1,6 +0,0 @@ -# Re-export from canonical location for backward compatibility. -# New code should import from `services.workers.ray_utils` directly. -from services.workers.ray_utils import ( # noqa: F401 - call_ray_actor_with_timeout, - retry_with_backoff, -) diff --git a/openrag/components/reranker.py b/openrag/components/reranker.py deleted file mode 100644 index 14403732a..000000000 --- a/openrag/components/reranker.py +++ /dev/null @@ -1,82 +0,0 @@ -import asyncio - -from infinity_client import Client -from infinity_client.api.default import rerank -from infinity_client.models import RerankInput, ReRankResult -from langchain_core.documents.base import Document - - -class BaseReranker: - async def rerank(self, query: str, documents: list[Document], top_k: int | None = None) -> list[Document]: - """Rerank a list of documents based on a query and an optional top_k parameter""" - raise NotImplementedError("Rerank method must be implemented by subclasses") - - @staticmethod - def rrf_reranking(doc_lists: list[list], k: int = 60) -> list[Document]: - """Reciprocal_rank_fusion that takes multiple lists of ranked documents - and an optional parameter k used in the RRF formula - RRF formula: \\sum_{i=1}^{n} \frac{1}{k + rank_i} - where rank_i is the rank of the document in the i-th list and n is the number of lists. - - k small: High sensitivity to top ranks - k large: More balanced sensitivity across ranks - k = 60 a common and balanced choice in practice. - """ - - if len(doc_lists) == 1: - return doc_lists[0] - - fused_scores = {} - for doc_list in doc_lists: - doc_list: list[Document] - for rank, doc in enumerate(doc_list, start=1): - doc_id = doc.metadata.get("_id") - - score, d = fused_scores.get(doc_id, (0, doc)) - fused_scores[doc_id] = (score + (1 / (rank + k)), d) - - # sort the docs - reranked_docs = [doc for score, doc in sorted(fused_scores.values(), key=lambda x: x[0], reverse=True)] - return reranked_docs - - -class Reranker(BaseReranker): - def __init__(self, logger, config): - self.model_name = config.reranker.model_name - self.client = Client(base_url=config.reranker.base_url) - self.logger = logger - self.semaphore = asyncio.Semaphore(5) # Only allow 5 reranking operation at a time - self.temporal_reranking = config.reranker.get("temporal_reranking", False) - self.logger.debug("Reranker initialized", model_name=self.model_name) - - async def rerank(self, query: str, documents: list[Document], top_k: int | None = None) -> list[Document]: - async with self.semaphore: - self.logger.debug("Reranking documents", documents_count=len(documents), top_k=top_k) - top_k = min(top_k, len(documents)) if top_k is not None else len(documents) - rerank_input = RerankInput.from_dict( - { - "model": self.model_name, - "query": query, - "documents": [doc.page_content for doc in documents], - "top_n": top_k, - "return_documents": True, - "raw_scores": True, # Normalized score between 0 and 1 - } - ) - try: - rerank_result: ReRankResult = await rerank.asyncio(client=self.client, body=rerank_input) - output = [] - for rerank_res in rerank_result.results: - doc = documents[rerank_res.index] - doc.metadata["relevance_score"] = rerank_res.relevance_score - output.append(doc) - return output - - except Exception as e: - self.logger.error( - "Reranking failed", - error=str(e), - model_name=self.model_name, - documents_count=len(documents), - ) - raise e diff --git a/openrag/components/reranker/__init__.py b/openrag/components/reranker/__init__.py deleted file mode 100644 index 2b9742509..000000000 --- a/openrag/components/reranker/__init__.py +++ /dev/null @@ -1,43 +0,0 @@ -import asyncio - -import services.inference.reranker_clients # noqa: F401 — registers "infinity"/"openai" -from core.config.retrieval import RerankerConfig -from core.rerankers import reranker_registry - -from .base import BaseReranker - - -class _RerankerShim(BaseReranker): - """Wraps a core ``Reranker`` (str-in / (idx, score)-out) behind the - legacy ``BaseReranker`` interface (Document-in / Document-out).""" - - def __init__(self, delegate, semaphore: int = 3): - self._delegate = delegate - self._semaphore = asyncio.Semaphore(semaphore) - - async def rerank(self, query, documents, top_k=None): - async with self._semaphore: - texts = [doc.page_content for doc in documents] - ranked = await self._delegate.rerank(query, texts, top_k=top_k) - output = [] - for index, score in ranked: - if not 0 <= index < len(documents): - continue - doc = documents[index] - doc.metadata["relevance_score"] = score - output.append(doc) - return output - - -class RerankerFactory: - @staticmethod - def get_reranker(reranker_config: RerankerConfig) -> BaseReranker: - provider = reranker_config.provider - delegate = reranker_registry.create( - provider, - endpoint=reranker_config.base_url, - model_name=reranker_config.model_name, - api_key=reranker_config.api_key, - timeout=reranker_config.timeout, - ) - return _RerankerShim(delegate, semaphore=reranker_config.semaphore) diff --git a/openrag/components/reranker/base.py b/openrag/components/reranker/base.py deleted file mode 100644 index 19a3defa3..000000000 --- a/openrag/components/reranker/base.py +++ /dev/null @@ -1,18 +0,0 @@ -from abc import ABC, abstractmethod - -from core.retrieval.rrf import rrf_reranking -from langchain_core.documents.base import Document - - -class BaseReranker(ABC): - @abstractmethod - async def rerank(self, query: str, documents: list[Document], top_k: int | None = None) -> list[Document]: - """Rerank a list of documents based on a query and an optional top_k parameter""" - - @staticmethod - def rrf_reranking(doc_lists: list[list[Document]], k: int = 60) -> list[Document]: - return rrf_reranking( - doc_lists, - key_fn=lambda doc: doc.metadata.get("_id", id(doc)), - k=k, - ) diff --git a/openrag/components/reranker/infinity.py b/openrag/components/reranker/infinity.py deleted file mode 100644 index 6135c4bec..000000000 --- a/openrag/components/reranker/infinity.py +++ /dev/null @@ -1,63 +0,0 @@ -"""Backward-compatibility shim — delegates to services.inference.reranker_clients. - -All new code should import directly from ``services.inference.reranker_clients``. -""" - -import asyncio - -from core.utils.logging import get_logger -from infinity_client import Client -from infinity_client.api.default import rerank -from infinity_client.models import RerankInput, ReRankResult -from langchain_core.documents.base import Document -from services.inference.reranker_clients import InfinityReranker as InfinityRerankerAdapter # noqa: F401 - -from .base import BaseReranker - -logger = get_logger() - - -class InfinityReranker(BaseReranker): - """Legacy InfinityReranker. New code should use InfinityRerankerAdapter (via DI).""" - - def __init__(self, config): - self.model_name = config.reranker.model_name - self.client = Client( - base_url=config.reranker.base_url, - timeout=config.reranker.timeout, - headers={"Authorization": f"Bearer {config.reranker.api_key}"}, - ) - self.semaphore = asyncio.Semaphore(config.reranker.semaphore) - logger.debug("Reranker initialized", model_name=self.model_name) - - async def rerank(self, query: str, documents: list[Document], top_k: int | None = None) -> list[Document]: - async with self.semaphore: - logger.debug("Reranking documents", documents_count=len(documents), top_k=top_k) - top_k = min(top_k, len(documents)) if top_k is not None else len(documents) - rerank_input = RerankInput.from_dict( - { - "model": self.model_name, - "query": query, - "documents": [doc.page_content for doc in documents], - "top_n": top_k, - "return_documents": True, - "raw_scores": True, - } - ) - try: - rerank_result: ReRankResult = await rerank.asyncio(client=self.client, body=rerank_input) - output = [] - for rerank_res in rerank_result.results: - doc = documents[rerank_res.index] - doc.metadata["relevance_score"] = rerank_res.relevance_score - output.append(doc) - return output - - except Exception as e: - logger.error( - "Reranking failed", - error=str(e), - model_name=self.model_name, - documents_count=len(documents), - ) - return documents[:top_k] diff --git a/openrag/components/reranker/openai.py b/openrag/components/reranker/openai.py deleted file mode 100644 index 9af3c5947..000000000 --- a/openrag/components/reranker/openai.py +++ /dev/null @@ -1,64 +0,0 @@ -"""Backward-compatibility shim — delegates to services.inference.reranker_clients. - -All new code should import directly from ``services.inference.reranker_clients``. -""" - -import asyncio - -import httpx -from core.utils.logging import get_logger -from langchain_core.documents.base import Document -from services.inference.reranker_clients import OpenAIReranker as OpenAIRerankerAdapter # noqa: F401 - -from .base import BaseReranker - -logger = get_logger() - - -class OpenAIReranker(BaseReranker): - """Legacy OpenAIReranker. New code should use OpenAIRerankerAdapter (via DI).""" - - def __init__(self, config): - self.model_name = config.reranker.model_name - base_url = config.reranker.base_url.rstrip("/") - self.rerank_url = f"{base_url}/rerank" - self.semaphore = asyncio.Semaphore(config.reranker.semaphore) - self.timeout = config.reranker.timeout - self.client = httpx.AsyncClient( - headers={"Authorization": f"Bearer {config.reranker.api_key}"}, - ) - logger.debug("OpenAI Reranker initialized", model_name=self.model_name) - - async def rerank(self, query: str, documents: list[Document], top_k: int | None = None) -> list[Document]: - async with self.semaphore: - logger.debug("Reranking documents", documents_count=len(documents), top_k=top_k) - top_k = min(top_k, len(documents)) if top_k is not None else len(documents) - try: - response = await self.client.post( - self.rerank_url, - json={ - "model": self.model_name, - "query": query, - "documents": [doc.page_content for doc in documents], - "top_n": top_k, - }, - timeout=self.timeout, - ) - response.raise_for_status() - data = response.json() - - output = [] - for result in data["results"]: - doc = documents[result["index"]] - doc.metadata["relevance_score"] = result["relevance_score"] - output.append(doc) - return output - - except Exception as e: - logger.error( - "Reranking failed", - error=str(e), - model_name=self.model_name, - documents_count=len(documents), - ) - return documents[:top_k] diff --git a/openrag/components/reranker/test_rrf_reranking.py b/openrag/components/reranker/test_rrf_reranking.py deleted file mode 100644 index 40efb456d..000000000 --- a/openrag/components/reranker/test_rrf_reranking.py +++ /dev/null @@ -1,75 +0,0 @@ -"""Tests for BaseReranker.rrf_reranking static method.""" - -from langchain_core.documents.base import Document - -from .base import BaseReranker - - -def make_doc(doc_id: str, content: str = "", **metadata) -> Document: - return Document(page_content=content, metadata={"_id": doc_id, **metadata}) - - -class TestRrfRerankingSingleList: - def test_single_list_returned_as_list_copy(self): - docs = [make_doc("a"), make_doc("b"), make_doc("c")] - result = BaseReranker.rrf_reranking([docs]) - assert result == docs - assert result is not docs - - -class TestRrfRerankingMultipleLists: - def test_two_lists_no_overlap_all_docs_present(self): - list1 = [make_doc("a"), make_doc("b")] - list2 = [make_doc("c"), make_doc("d")] - result = BaseReranker.rrf_reranking([list1, list2]) - assert {d.metadata["_id"] for d in result} == {"a", "b", "c", "d"} - - def test_document_in_multiple_lists_ranked_higher(self): - # doc_shared appears rank 1 in both lists - # doc_only_list1 appears rank 2 in list1 only - doc_shared = make_doc("shared") - doc_only = make_doc("only") - list1 = [doc_shared, doc_only] - list2 = [doc_shared] - result = BaseReranker.rrf_reranking([list1, list2]) - ids = [d.metadata["_id"] for d in result] - assert ids[0] == "shared" - - def test_overlapping_docs_deduped(self): - doc = make_doc("dup") - list1 = [doc, make_doc("a")] - list2 = [doc, make_doc("b")] - result = BaseReranker.rrf_reranking([list1, list2]) - ids = [d.metadata["_id"] for d in result] - assert ids.count("dup") == 1 - - def test_sorted_descending_by_rrf_score(self): - # doc_a: rank 1 in both → score = 2/(1+60) ≈ 0.0328 - # doc_b: rank 2 in both → score = 2/(2+60) ≈ 0.0323 - # doc_c: rank 3 in both → score = 2/(3+60) ≈ 0.0317 - doc_a, doc_b, doc_c = make_doc("a"), make_doc("b"), make_doc("c") - list1 = [doc_a, doc_b, doc_c] - list2 = [doc_a, doc_b, doc_c] - result = BaseReranker.rrf_reranking([list1, list2]) - assert [d.metadata["_id"] for d in result] == ["a", "b", "c"] - - def test_higher_rank_in_one_list_can_outweigh_single_appearance(self): - # doc_top: rank 1 in list1 only → score = 1/61 - # doc_bottom: rank 3 in list1, rank 3 in list2 → score = 2/63 - # 2/63 ≈ 0.0317 > 1/61 ≈ 0.0164, so doc_bottom wins - doc_top = make_doc("top") - doc_bottom = make_doc("bottom") - list1 = [doc_top, make_doc("x"), doc_bottom] - list2 = [make_doc("y"), make_doc("z"), doc_bottom] - result = BaseReranker.rrf_reranking([list1, list2]) - ids = [d.metadata["_id"] for d in result] - assert ids.index("bottom") < ids.index("top") - - -class TestRrfRerankingDocumentPreservation: - def test_preserves_page_content_and_metadata(self): - doc = make_doc("doc1", content="hello world", source="file.pdf", score=0.9) - result = BaseReranker.rrf_reranking([[doc], [doc]]) - assert result[0].page_content == "hello world" - assert result[0].metadata["source"] == "file.pdf" - assert result[0].metadata["score"] == 0.9 diff --git a/openrag/components/test_files.py b/openrag/components/test_files.py deleted file mode 100644 index eb3523f67..000000000 --- a/openrag/components/test_files.py +++ /dev/null @@ -1,55 +0,0 @@ -import io -from pathlib import Path - -import pytest -from components.files import save_file_to_disk -from fastapi import UploadFile - - -@pytest.mark.asyncio -async def test_save_file_to_disk_writes_content(tmp_path: Path): - content = b"hello world" - upload = UploadFile( - file=io.BytesIO(content), - filename="test.bin", - ) - - dest_dir = tmp_path / "uploads" - - saved_path = await save_file_to_disk(file=upload, dest_dir=dest_dir, chunk_size=4) - - assert saved_path.exists() - assert saved_path.parent == dest_dir - assert saved_path.name == "test.bin" - - with open(saved_path, "rb") as f: - saved_content = f.read() - - assert saved_content == content - - -@pytest.mark.asyncio -async def test_save_file_to_disk_with_random_prefix(tmp_path, monkeypatch): - def fake_make_unique_filename(filename: str) -> str: - assert filename == "test.txt" - return "PREFIX_1234_test.txt" - - monkeypatch.setattr("components.files.make_unique_filename", fake_make_unique_filename) - - file_content = b"hello world" - upload = UploadFile( - filename="test.txt", - file=io.BytesIO(file_content), - ) - - saved_path = await save_file_to_disk( - file=upload, - dest_dir=tmp_path, - chunk_size=1024, - with_random_prefix=True, - ) - - assert saved_path.parent == tmp_path - assert saved_path.name == "PREFIX_1234_test.txt" - assert saved_path.exists() - assert saved_path.read_bytes() == file_content diff --git a/openrag/components/test_llm.py b/openrag/components/test_llm.py deleted file mode 100644 index 822313be3..000000000 --- a/openrag/components/test_llm.py +++ /dev/null @@ -1,97 +0,0 @@ -import pytest -from components.llm import LLM -from config.models import LLMConfig - - -@pytest.fixture -def llm(): - return LLM( - LLMConfig( - base_url="http://default-llm:8000/v1", - api_key="default-key", - model="default-model", - temperature=0.3, - ) - ) - - -class TestExtractLlmOverrides: - def test_no_override_uses_defaults(self, llm): - request = { - "model": "openrag-my-partition", - "messages": [{"role": "user", "content": "hello"}], - "stream": False, - } - payload, base_url, headers = llm._extract_llm_overrides(request) - - assert payload["model"] == "default-model" - assert payload["temperature"] == 0.3 - assert base_url == "http://default-llm:8000/v1" - assert headers["Authorization"] == "Bearer default-key" - - def test_override_all_fields(self, llm): - request = { - "model": "openrag-my-partition", - "messages": [{"role": "user", "content": "hello"}], - "stream": False, - "metadata": { - "llm_override": { - "base_url": "http://custom-llm:9000/v1", - "api_key": "custom-key", - "model": "custom-model", - } - }, - } - payload, base_url, headers = llm._extract_llm_overrides(request) - - assert payload["model"] == "custom-model" - assert base_url == "http://custom-llm:9000/v1" - assert headers["Authorization"] == "Bearer custom-key" - - def test_trailing_slash_stripped_from_base_url(self, llm): - request = { - "model": "openrag-my-partition", - "stream": False, - "metadata": {"llm_override": {"base_url": "http://custom:8000/v1///"}}, - } - _, base_url, _ = llm._extract_llm_overrides(request) - - assert base_url == "http://custom:8000/v1" - - def test_request_params_forwarded_to_payload(self, llm): - request = { - "model": "openrag-my-partition", - "messages": [{"role": "user", "content": "hello"}], - "stream": True, - "max_tokens": 2048, - "temperature": 0.9, - } - payload, _, _ = llm._extract_llm_overrides(request) - - assert payload["stream"] is True - assert payload["max_tokens"] == 2048 - assert payload["temperature"] == 0.9 - assert payload["messages"] == [{"role": "user", "content": "hello"}] - - def test_metadata_without_llm_override_uses_defaults(self, llm): - request = { - "model": "openrag-my-partition", - "stream": False, - "metadata": {"use_map_reduce": True}, - } - payload, base_url, headers = llm._extract_llm_overrides(request) - - assert payload["model"] == "default-model" - assert base_url == "http://default-llm:8000/v1" - assert headers["Authorization"] == "Bearer default-key" - - def test_llm_override_popped_from_metadata(self, llm): - metadata = { - "use_map_reduce": False, - "llm_override": {"model": "custom"}, - } - request = {"model": "x", "stream": False, "metadata": metadata} - llm._extract_llm_overrides(request) - - assert "llm_override" not in metadata - assert "use_map_reduce" in metadata diff --git a/openrag/components/test_relationships.py b/openrag/components/test_relationships.py deleted file mode 100644 index a404a7014..000000000 --- a/openrag/components/test_relationships.py +++ /dev/null @@ -1,669 +0,0 @@ -""" -Unit tests for document relationship functionality. - -Tests the relationship_id and parent_id fields for linking related documents -(e.g., email threads, folder hierarchies). - -Note: These tests use an in-memory SQLite database to test the PartitionFileManager -methods without requiring the full application stack. -""" - -import json - -import pytest -from sqlalchemy import Column, DateTime, Index, Integer, String, Text, create_engine, text -from sqlalchemy.orm import declarative_base, sessionmaker -from sqlalchemy.sql import func - -# Create isolated SQLAlchemy base for testing -TestBase = declarative_base() - - -class FileModel(TestBase): - """Test version of File model with relationship fields.""" - - __tablename__ = "files" - - id = Column(Integer, primary_key=True, autoincrement=True) - file_id = Column(String, nullable=False) - partition_name = Column(String, nullable=False) - file_metadata = Column(Text, nullable=True) - created_at = Column(DateTime(timezone=True), server_default=func.now()) - updated_at = Column(DateTime(timezone=True), onupdate=func.now()) - relationship_id = Column(String, nullable=True, index=True) - parent_id = Column(String, nullable=True, index=True) - - __table_args__ = ( - Index("ix_relationship_partition", "relationship_id", "partition_name"), - Index("ix_parent_partition", "parent_id", "partition_name"), - ) - - def to_dict(self): - return { - "file_id": self.file_id, - "partition": self.partition_name, - "file_metadata": (json.loads(self.file_metadata) if self.file_metadata else {}), - "relationship_id": self.relationship_id, - "parent_id": self.parent_id, - "created_at": str(self.created_at) if self.created_at else None, - } - - -class PartitionFileManagerHelper: - """Test version of PartitionFileManager for isolated testing.""" - - def __init__(self, session_factory): - self.Session = session_factory - - def add_file_to_partition( - self, - partition: str, - file_id: str, - file_metadata: dict = None, - relationship_id: str = None, - parent_id: str = None, - ): - """Add a file record to the database.""" - with self.Session() as session: - file_entry = FileModel( - file_id=file_id, - partition_name=partition, - file_metadata=json.dumps(file_metadata) if file_metadata else None, - relationship_id=relationship_id, - parent_id=parent_id, - ) - session.add(file_entry) - session.commit() - - def get_files_by_relationship(self, partition: str, relationship_id: str) -> list[dict]: - """Get all files with the same relationship_id in a partition.""" - with self.Session() as session: - files = ( - session.query(FileModel) - .filter( - FileModel.partition_name == partition, - FileModel.relationship_id == relationship_id, - ) - .all() - ) - return [f.to_dict() for f in files] - - def get_file_ids_by_relationship(self, partition: str, relationship_id: str) -> list[str]: - """Get file IDs for files with the same relationship_id.""" - with self.Session() as session: - files = ( - session.query(FileModel.file_id) - .filter( - FileModel.partition_name == partition, - FileModel.relationship_id == relationship_id, - ) - .all() - ) - return [f.file_id for f in files] - - def get_file_ancestors(self, partition: str, file_id: str, max_ancestor_depth: int | None = None) -> list[dict]: - """Get the ancestor chain for a file using recursive CTE.""" - with self.Session() as session: - # Recursive CTE for ancestor traversal with optional max depth - depth_condition = "WHERE a.depth < :max_ancestor_depth" if max_ancestor_depth is not None else "" - query = text(f""" - WITH RECURSIVE ancestors AS ( - -- Base case: start with the target file - SELECT id, file_id, partition_name, parent_id, file_metadata, - relationship_id, 0 as depth - FROM files - WHERE file_id = :file_id AND partition_name = :partition - AND relationship_id IS NOT NULL - - UNION ALL - - -- Recursive case: get parent - SELECT f.id, f.file_id, f.partition_name, f.parent_id, - f.file_metadata, f.relationship_id, a.depth + 1 - FROM files f - INNER JOIN ancestors a ON f.file_id = a.parent_id - AND f.partition_name = a.partition_name - AND f.relationship_id IS NOT NULL - {depth_condition} - ) - SELECT * FROM ancestors ORDER BY depth DESC - """) - - params = {"file_id": file_id, "partition": partition} - if max_ancestor_depth is not None: - params["max_ancestor_depth"] = max_ancestor_depth - - result = session.execute(query, params) - - rows = result.fetchall() - return [ - { - "file_id": row.file_id, - "partition": row.partition_name, - "file_metadata": (json.loads(row.file_metadata) if row.file_metadata else {}), - "relationship_id": row.relationship_id, - "parent_id": row.parent_id, - } - for row in rows - ] - - def get_ancestor_file_ids(self, partition: str, file_id: str, max_ancestor_depth: int | None = None) -> list[str]: - """Get file IDs of ancestors.""" - ancestors = self.get_file_ancestors(partition, file_id, max_ancestor_depth=max_ancestor_depth) - return [a["file_id"] for a in ancestors] - - -@pytest.fixture -def in_memory_db(): - """Create an in-memory SQLite database for testing.""" - engine = create_engine("sqlite:///:memory:") - TestBase.metadata.create_all(engine) - Session = sessionmaker(bind=engine) - return Session - - -@pytest.fixture -def file_manager(in_memory_db): - """Create a PartitionFileManagerHelper with in-memory database.""" - return PartitionFileManagerHelper(in_memory_db) - - -class TestAddFileWithRelationships: - """Test adding files with relationship_id and parent_id.""" - - def test_add_file_with_relationship_id(self, file_manager): - """Test that files can be added with a relationship_id.""" - file_manager.add_file_to_partition( - partition="test_partition", - file_id="file_001", - file_metadata={"filename": "email1.eml"}, - relationship_id="thread_abc123", - ) - - with file_manager.Session() as session: - result = session.execute(text("SELECT relationship_id FROM files WHERE file_id = 'file_001'")).fetchone() - assert result[0] == "thread_abc123" - - def test_add_file_with_parent_id(self, file_manager): - """Test that files can be added with a parent_id.""" - file_manager.add_file_to_partition( - partition="test_partition", - file_id="file_001", - file_metadata={"filename": "reply.eml"}, - parent_id="file_000", - ) - - with file_manager.Session() as session: - result = session.execute(text("SELECT parent_id FROM files WHERE file_id = 'file_001'")).fetchone() - assert result[0] == "file_000" - - def test_add_file_with_both_relationship_and_parent(self, file_manager): - """Test that files can be added with both relationship_id and parent_id.""" - file_manager.add_file_to_partition( - partition="test_partition", - file_id="file_002", - file_metadata={"filename": "reply2.eml"}, - relationship_id="thread_abc123", - parent_id="file_001", - ) - - with file_manager.Session() as session: - result = session.execute( - text("SELECT relationship_id, parent_id FROM files WHERE file_id = 'file_002'") - ).fetchone() - assert result[0] == "thread_abc123" - assert result[1] == "file_001" - - def test_add_file_without_relationships(self, file_manager): - """Test that files can be added without relationship fields (backward compat).""" - file_manager.add_file_to_partition( - partition="test_partition", - file_id="file_003", - file_metadata={"filename": "standalone.pdf"}, - ) - - with file_manager.Session() as session: - result = session.execute( - text("SELECT relationship_id, parent_id FROM files WHERE file_id = 'file_003'") - ).fetchone() - assert result[0] is None - assert result[1] is None - - -class TestGetFilesByRelationship: - """Test querying files by relationship_id.""" - - def test_get_files_by_relationship(self, file_manager): - """Test retrieving all files with the same relationship_id.""" - # Add multiple files with same relationship_id - for i in range(3): - file_manager.add_file_to_partition( - partition="test_partition", - file_id=f"email_{i}", - file_metadata={"filename": f"email{i}.eml"}, - relationship_id="thread_xyz", - ) - - # Add a file with different relationship_id - file_manager.add_file_to_partition( - partition="test_partition", - file_id="email_other", - file_metadata={"filename": "other.eml"}, - relationship_id="thread_other", - ) - - results = file_manager.get_files_by_relationship( - partition="test_partition", - relationship_id="thread_xyz", - ) - - assert len(results) == 3 - file_ids = [r["file_id"] for r in results] - assert set(file_ids) == {"email_0", "email_1", "email_2"} - - def test_get_files_by_relationship_empty_result(self, file_manager): - """Test that empty list is returned when no files match.""" - file_manager.add_file_to_partition( - partition="test_partition", - file_id="file_001", - file_metadata={"filename": "test.pdf"}, - relationship_id="rel_abc", - ) - - results = file_manager.get_files_by_relationship( - partition="test_partition", - relationship_id="nonexistent", - ) - - assert results == [] - - def test_get_files_by_relationship_respects_partition(self, file_manager): - """Test that relationship query respects partition boundaries.""" - # Add files with same relationship_id in different partitions - file_manager.add_file_to_partition( - partition="partition_a", - file_id="file_a", - file_metadata={"filename": "a.pdf"}, - relationship_id="shared_rel", - ) - file_manager.add_file_to_partition( - partition="partition_b", - file_id="file_b", - file_metadata={"filename": "b.pdf"}, - relationship_id="shared_rel", - ) - - results = file_manager.get_files_by_relationship( - partition="partition_a", - relationship_id="shared_rel", - ) - - assert len(results) == 1 - assert results[0]["file_id"] == "file_a" - - def test_get_file_ids_by_relationship(self, file_manager): - """Test retrieving only file IDs by relationship_id.""" - for i in range(3): - file_manager.add_file_to_partition( - partition="test_partition", - file_id=f"doc_{i}", - file_metadata={"filename": f"doc{i}.pdf"}, - relationship_id="folder_123", - ) - - file_ids = file_manager.get_file_ids_by_relationship( - partition="test_partition", - relationship_id="folder_123", - ) - - assert len(file_ids) == 3 - assert set(file_ids) == {"doc_0", "doc_1", "doc_2"} - - -class TestGetFileAncestors: - """Test retrieving ancestor chain for a file.""" - - def test_get_file_ancestors_single_file(self, file_manager): - """Test that a file with no parent but with a relationship_id returns only itself.""" - file_manager.add_file_to_partition( - partition="test_partition", - file_id="root_email", - file_metadata={"filename": "root.eml"}, - relationship_id="thread_single", - ) - - ancestors = file_manager.get_file_ancestors( - partition="test_partition", - file_id="root_email", - ) - - assert len(ancestors) == 1 - assert ancestors[0]["file_id"] == "root_email" - - def test_get_file_ancestors_chain(self, file_manager): - """Test retrieving a chain of ancestors.""" - # Create email thread: root -> reply1 -> reply2 - file_manager.add_file_to_partition( - partition="test_partition", - file_id="email_root", - file_metadata={"filename": "root.eml"}, - relationship_id="thread_1", - ) - file_manager.add_file_to_partition( - partition="test_partition", - file_id="email_reply1", - file_metadata={"filename": "reply1.eml"}, - relationship_id="thread_1", - parent_id="email_root", - ) - file_manager.add_file_to_partition( - partition="test_partition", - file_id="email_reply2", - file_metadata={"filename": "reply2.eml"}, - relationship_id="thread_1", - parent_id="email_reply1", - ) - - ancestors = file_manager.get_file_ancestors( - partition="test_partition", - file_id="email_reply2", - ) - - # Should return [root, reply1, reply2] in order from root to target - assert len(ancestors) == 3 - file_ids = [a["file_id"] for a in ancestors] - assert file_ids == ["email_root", "email_reply1", "email_reply2"] - - def test_get_file_ancestors_returns_ordered_path(self, file_manager): - """Test that ancestors are returned in correct order (root first).""" - # Create deeper hierarchy: A -> B -> C -> D - file_manager.add_file_to_partition( - partition="test_partition", - file_id="file_a", - file_metadata={"filename": "a.txt"}, - relationship_id="thread_ordered", - ) - file_manager.add_file_to_partition( - partition="test_partition", - file_id="file_b", - file_metadata={"filename": "b.txt"}, - parent_id="file_a", - relationship_id="thread_ordered", - ) - file_manager.add_file_to_partition( - partition="test_partition", - file_id="file_c", - file_metadata={"filename": "c.txt"}, - parent_id="file_b", - relationship_id="thread_ordered", - ) - file_manager.add_file_to_partition( - partition="test_partition", - file_id="file_d", - file_metadata={"filename": "d.txt"}, - parent_id="file_c", - relationship_id="thread_ordered", - ) - - ancestors = file_manager.get_file_ancestors( - partition="test_partition", - file_id="file_d", - ) - - # Verify order: root first, target last - assert len(ancestors) == 4 - file_ids = [a["file_id"] for a in ancestors] - assert file_ids == ["file_a", "file_b", "file_c", "file_d"] - - def test_get_file_ancestors_nonexistent_file(self, file_manager): - """Test that empty list is returned for nonexistent file.""" - ancestors = file_manager.get_file_ancestors( - partition="test_partition", - file_id="nonexistent", - ) - - assert ancestors == [] - - def test_get_ancestor_file_ids(self, file_manager): - """Test retrieving only ancestor file IDs.""" - # Create chain: root -> child - file_manager.add_file_to_partition( - partition="test_partition", - file_id="parent_file", - file_metadata={"filename": "parent.txt"}, - relationship_id="thread_ids", - ) - file_manager.add_file_to_partition( - partition="test_partition", - file_id="child_file", - file_metadata={"filename": "child.txt"}, - parent_id="parent_file", - relationship_id="thread_ids", - ) - - ancestor_ids = file_manager.get_ancestor_file_ids( - partition="test_partition", - file_id="child_file", - ) - - assert len(ancestor_ids) == 2 - assert ancestor_ids == ["parent_file", "child_file"] - - def test_get_file_ancestors_max_ancestor_depth_none_returns_all(self, file_manager): - """Test that max_ancestor_depth=None returns all ancestors (unlimited traversal).""" - # Create deep hierarchy: 0 -> 1 -> 2 -> ... -> 5 - file_manager.add_file_to_partition( - partition="test_partition", - file_id="level_0", - file_metadata={"filename": "root.txt"}, - relationship_id="thread_depth_none", - ) - for i in range(1, 6): - file_manager.add_file_to_partition( - partition="test_partition", - file_id=f"level_{i}", - file_metadata={"filename": f"level_{i}.txt"}, - parent_id=f"level_{i - 1}", - relationship_id="thread_depth_none", - ) - - # Without max_ancestor_depth (None), should return all 6 levels - ancestors = file_manager.get_file_ancestors( - partition="test_partition", - file_id="level_5", - max_ancestor_depth=None, - ) - - assert len(ancestors) == 6 - file_ids = [a["file_id"] for a in ancestors] - assert file_ids == ["level_0", "level_1", "level_2", "level_3", "level_4", "level_5"] - - def test_get_file_ancestors_max_ancestor_depth_limits_traversal(self, file_manager): - """Test that max_ancestor_depth limits how many ancestors are returned.""" - # Create deep hierarchy: 0 -> 1 -> 2 -> ... -> 5 - file_manager.add_file_to_partition( - partition="test_partition", - file_id="node_0", - file_metadata={"filename": "root.txt"}, - relationship_id="thread_depth_limit", - ) - for i in range(1, 6): - file_manager.add_file_to_partition( - partition="test_partition", - file_id=f"node_{i}", - file_metadata={"filename": f"node_{i}.txt"}, - parent_id=f"node_{i - 1}", - relationship_id="thread_depth_limit", - ) - - # With max_ancestor_depth=2, should return target (depth 0) + 2 ancestors - ancestors = file_manager.get_file_ancestors( - partition="test_partition", - file_id="node_5", - max_ancestor_depth=2, - ) - - # Should get node_5 (depth 0), node_4 (depth 1), node_3 (depth 2) - assert len(ancestors) == 3 - file_ids = [a["file_id"] for a in ancestors] - assert file_ids == ["node_3", "node_4", "node_5"] - - def test_get_file_ancestors_max_ancestor_depth_zero_returns_only_target(self, file_manager): - """Test that max_ancestor_depth=0 returns only the target file itself.""" - file_manager.add_file_to_partition( - partition="test_partition", - file_id="root", - file_metadata={"filename": "root.txt"}, - relationship_id="thread_depth_zero", - ) - file_manager.add_file_to_partition( - partition="test_partition", - file_id="child", - file_metadata={"filename": "child.txt"}, - parent_id="root", - relationship_id="thread_depth_zero", - ) - - # max_ancestor_depth=0 means no traversal beyond the target - ancestors = file_manager.get_file_ancestors( - partition="test_partition", - file_id="child", - max_ancestor_depth=0, - ) - - # Should only return the target file (depth 0 is included, but no recursion) - assert len(ancestors) == 1 - assert ancestors[0]["file_id"] == "child" - - def test_get_file_ancestors_max_ancestor_depth_exceeds_chain_length(self, file_manager): - """Test that max_ancestor_depth larger than chain length returns full chain.""" - # Create short chain: A -> B -> C - file_manager.add_file_to_partition( - partition="test_partition", - file_id="short_0", - file_metadata={"filename": "a.txt"}, - relationship_id="thread_short", - ) - file_manager.add_file_to_partition( - partition="test_partition", - file_id="short_1", - file_metadata={"filename": "b.txt"}, - parent_id="short_0", - relationship_id="thread_short", - ) - file_manager.add_file_to_partition( - partition="test_partition", - file_id="short_2", - file_metadata={"filename": "c.txt"}, - parent_id="short_1", - relationship_id="thread_short", - ) - - # max_ancestor_depth=100 but chain is only 3 levels - ancestors = file_manager.get_file_ancestors( - partition="test_partition", - file_id="short_2", - max_ancestor_depth=100, - ) - - # Should return all 3 levels - assert len(ancestors) == 3 - file_ids = [a["file_id"] for a in ancestors] - assert file_ids == ["short_0", "short_1", "short_2"] - - def test_get_ancestor_file_ids_with_max_ancestor_depth(self, file_manager): - """Test that get_ancestor_file_ids respects max_ancestor_depth parameter.""" - # Create chain: A -> B -> C -> D - file_manager.add_file_to_partition( - partition="test_partition", - file_id="chain_0", - file_metadata={"filename": "a.txt"}, - relationship_id="thread_chain", - ) - for i in range(1, 4): - file_manager.add_file_to_partition( - partition="test_partition", - file_id=f"chain_{i}", - file_metadata={"filename": f"{chr(97 + i)}.txt"}, - parent_id=f"chain_{i - 1}", - relationship_id="thread_chain", - ) - - # With max_ancestor_depth=1, should get target + 1 ancestor - ancestor_ids = file_manager.get_ancestor_file_ids( - partition="test_partition", - file_id="chain_3", - max_ancestor_depth=1, - ) - - assert len(ancestor_ids) == 2 - assert ancestor_ids == ["chain_2", "chain_3"] - - -class TestStandaloneFileNoExpansion: - """Test that a file indexed without relationship_id yields no additional chunks - when include_related and include_ancestors are both active.""" - - def test_no_extra_chunks_for_file_without_relationship_id(self, file_manager): - """A standalone file (no relationship_id, no parent_id) must not bring - additional files when both include_related and include_ancestors are activated. - - Mirrors the logic in _expand_with_related_chunks: - - include_related: the guard `metadata.get("relationship_id")` is falsy, - so no related lookup is issued and the related task set stays empty. - - include_ancestors: get_file_ancestors returns only the file itself when - there is no parent, so it is already in seen_ids — nothing new is added. - """ - file_manager.add_file_to_partition( - partition="test_partition", - file_id="standalone", - file_metadata={"filename": "standalone.pdf"}, - # No relationship_id, no parent_id - ) - - # Verify the file has no relationship_id (the falsy guard that prevents - # the include_related lookup from being issued at all). - files = file_manager.get_files_by_relationship( - partition="test_partition", - relationship_id="standalone", # non-existent → empty - ) - assert files == [], "No files should share a relationship with a standalone file" - - with file_manager.Session() as session: - row = session.execute(text("SELECT relationship_id FROM files WHERE file_id = 'standalone'")).fetchone() - assert not row[0], "relationship_id must be falsy so include_related is skipped" - - ancestors = file_manager.get_file_ancestors( - partition="test_partition", - file_id="standalone", - ) - assert len(ancestors) == 0, ( - "Standalone file has no relationship_id, so ancestor list must be empty — " - "the relationship_id filter in get_file_ancestors excludes it" - ) - - -class TestFileModelFields: - """Test that File model correctly handles relationship fields.""" - - def test_to_dict_includes_relationship_fields(self, file_manager): - """Test that to_dict() includes relationship_id and parent_id.""" - file_manager.add_file_to_partition( - partition="test_partition", - file_id="file_001", - file_metadata={"filename": "test.eml", "subject": "Hello"}, - relationship_id="thread_abc", - parent_id="file_000", - ) - - files = file_manager.get_files_by_relationship( - partition="test_partition", - relationship_id="thread_abc", - ) - - assert len(files) == 1 - file_dict = files[0] - assert file_dict["relationship_id"] == "thread_abc" - assert file_dict["parent_id"] == "file_000" - assert file_dict["file_metadata"]["filename"] == "test.eml" - assert file_dict["file_metadata"]["subject"] == "Hello" diff --git a/openrag/components/utils.py b/openrag/components/utils.py deleted file mode 100644 index 4931ce226..000000000 --- a/openrag/components/utils.py +++ /dev/null @@ -1,326 +0,0 @@ -import asyncio -import copy -import json -import re -import threading -from typing import ClassVar - -from config import load_config -from core.utils.logging import get_logger -from fast_langdetect import LangDetectConfig, LangDetector -from langchain_core.documents.base import Document -from services.inference.distributed_semaphore import ( - DistributedSemaphore, # noqa: F401 - DistributedSemaphoreActor, # noqa: F401 -) - -SOURCE_SEPARATOR = "-" * 10 + "\n\n" - -logger = get_logger() - - -class SingletonMeta(type): - _instances: ClassVar[dict] = {} - _lock = threading.Lock() # Ensures thread safety - - def __call__(cls, *args, **kwargs): - if cls not in cls._instances: # First check (not thread-safe yet) - with cls._lock: # Prevents multiple threads from creating instances - if cls not in cls._instances: # Second check (double-checked locking) - instance = super().__call__(*args, **kwargs) - cls._instances[cls] = instance - return cls._instances[cls] - - -_cached_length_function = None - - -def get_num_tokens(): - global _cached_length_function - if _cached_length_function is None: - try: - from langchain_openai import ChatOpenAI - - config = load_config() - llm = ChatOpenAI(**config.llm.model_dump()) - _cached_length_function = llm.get_num_tokens - except Exception as exc: - # ChatOpenAI validates an openai client at construction, which - # requires a non-empty api_key. Token counting itself is local - # (tiktoken) and needs no key or network, so fall back to a - # tiktoken encoder when the client cannot be built (keyless - # deployments / CI mock-vLLM). cl100k_base matches the - # GPT-3.5/4 family OpenRAG targets; counts are equivalent. - import tiktoken - - logger.warning( - "ChatOpenAI unavailable for token counting, falling back to tiktoken cl100k_base", - error=str(exc), - ) - _encoding = tiktoken.get_encoding("cl100k_base") - _cached_length_function = lambda text: len(_encoding.encode(text)) # noqa: E731 - return _cached_length_function - - -def format_context( - docs: list[Document], max_context_tokens: int = 4096, number_sources: bool = True -) -> tuple[str, list[int]]: - """Backward-compat shim — delegates to `core.prompts.chat_prompt_builder.format_context`. - - The legacy signature took LangChain Documents and resolved a tokenizer - internally; the core version takes raw strings + an injected - length_function. We adapt by extracting page_content and threading - the cached tokenizer through. - """ - from core.prompts.chat_prompt_builder import format_context as _core_format_context - - texts = [doc.page_content for doc in docs] - text, included = _core_format_context( - texts, - max_context_tokens=max_context_tokens, - length_function=get_num_tokens(), - number_sources=number_sources, - ) - logger.debug("Context formatted", doc_count=len(included)) - return text, included - - -def format_web_context( - web_results: list, - start_index: int = 1, - max_tokens: int = 2000, -) -> tuple[str, list[int], int]: - """Backward-compat shim — delegates to `core.prompts.chat_prompt_builder.format_web_context`. - - Same adaptation pattern as `format_context`: legacy resolved the - tokenizer internally, core takes it as a parameter. - """ - from core.prompts.chat_prompt_builder import format_web_context as _core_format_web_context - - text, source_numbers, total_tokens = _core_format_web_context( - web_results, - length_function=get_num_tokens(), - start_index=start_index, - max_tokens=max_tokens, - ) - logger.debug("Web context formatted", total_tokens=total_tokens, source_count=len(source_numbers)) - return text, source_numbers, total_tokens - - -# Line-terminal anchor `(?=\n|$)` — matches only when the tag sits flush against -# a newline or end-of-string. Safe: in-prose tags like "use [Sources: 1, 3] at end" -# are followed by text, so they stay. The LLM always places misplaced tags at the -# end of a sentence/bullet/line, which is exactly what this catches. -_SOURCES_NONE_RE = re.compile( - r"\n?[ \t]*\[?Sources?\]?\s*:\s*\[?\s*none\s*\]?[.\s]*?(?=\n|$)", - re.IGNORECASE, -) -_SOURCES_NUMS_RE = re.compile(r"\n?[ \t]*\[?Sources?\]?\s*:\s*\[?([\d,\s]+)\]?[.\s]*?(?=\n|$)") - - -def _strip_sources_tags(text: str) -> tuple[str, set[int], bool]: - """Strip every line-terminal [Sources: ...] tag. Return (cleaned, cited_nums, saw_none).""" - cited: set[int] = set() - for m in _SOURCES_NUMS_RE.finditer(text): - cited.update(int(n.strip()) for n in m.group(1).split(",") if n.strip().isdigit()) - saw_none = bool(_SOURCES_NONE_RE.search(text)) - cleaned = _SOURCES_NUMS_RE.sub("", text) - cleaned = _SOURCES_NONE_RE.sub("", cleaned) - return cleaned, cited, saw_none - - -def extract_and_strip_sources_block(text: str) -> tuple[str, set[int] | None]: - """Strip every line-terminal [Sources: ...] tag and return merged citations. - - Returns: - (clean_text, citations) where citations is: - - set of ints: union of all cited source numbers across every tag occurrence - - empty set: LLM said [Sources: none] and no numeric citations elsewhere - - None: no sources tag found — text returned unchanged - """ - cleaned, citations, saw_none = _strip_sources_tags(text) - - if not citations and not saw_none: - tail = text[-150:] if len(text) > 150 else text - logger.debug("No [Sources: ...] tag found in LLM response", tail=repr(tail)) - return text, None - - cleaned = cleaned.rstrip() - if citations: - logger.debug("Extracted source citations from LLM response", citations=sorted(citations)) - return cleaned, citations - - logger.debug("LLM explicitly reported no sources used") - return cleaned, set() - - -def filter_sources_by_citations(sources: list, citations: set[int] | None) -> list: - """Keep only sources whose 1-based index was cited. - - - citations is None: LLM didn't include tag → fallback to all sources - - citations is empty set: LLM said [Sources: none] → return no sources - - citations has values: filter to cited sources only - """ - if citations is None: - return sources - if not citations: - return [] - filtered = [s for i, s in enumerate(sources, start=1) if i in citations] - return filtered if filtered else sources - - -def _min_sources_tag_buffer_size(n_sources: int) -> int: - """Pessimistic upper bound on the length of a ``[Sources: 1, 2, …, N]`` tag. - - Worst case: every source cited, plus the wrapping characters and a small - margin for whitespace / leading newline variations the prompt allows. - """ - if n_sources <= 0: - return 100 - digits_total = sum(len(str(i)) for i in range(1, n_sources + 1)) - separators = max(0, n_sources - 1) * 2 # ", " - wrapping = len("\n[Sources: ") + len("]") + 8 # margin for whitespace / trailing chars - return digits_total + separators + wrapping - - -_MIN_STREAM_LOOKAHEAD = 80 - - -async def stream_with_source_filtering( - llm_stream, - sources: list, - model_name: str, - buffer_size: int | None = None, -): - """Process an LLM SSE stream, stripping every line-terminal [Sources: ...] tag. - - Look-ahead window: keep the last ``buffer_size`` chars of `pending` - buffered and emit everything before that. The held-back tail guarantees no - in-flight tag can straddle the emit boundary, so streaming flows in - real-time with bounded content lag regardless of newline cadence. - On stream end, flush the tail (EOS-anchored regex catches a final tag - without a trailing \n) and emit a finish chunk carrying the filtered - source metadata. - - ``buffer_size`` defaults to the minimum size that can hold a tag citing - *every* available source. Callers can override it explicitly (e.g. tests). - - Yields SSE "data: ..." lines ready to forward to the client. - """ - if buffer_size is None: - buffer_size = max(_MIN_STREAM_LOOKAHEAD, _min_sources_tag_buffer_size(len(sources))) - pending = "" - emitted_len = 0 - chunk_template = None - last_finish_reason = None - - async for line in llm_stream: - if not line.startswith("data:"): - continue - - if line.strip() == "data: [DONE]": - final_clean, citations = extract_and_strip_sources_block(pending) - final_clean = final_clean.rstrip() - - filtered = filter_sources_by_citations(sources, citations) - filtered_json = json.dumps({"sources": filtered}) - - if chunk_template and len(final_clean) > emitted_len: - tail_chunk = copy.deepcopy(chunk_template) - tail_chunk["choices"][0]["delta"] = {"content": final_clean[emitted_len:]} - tail_chunk["extra"] = filtered_json - yield f"data: {json.dumps(tail_chunk)}\n\n" - - if chunk_template: - # FIXME: race condition where clients miss sources because finish_reason - # arrives before the sources metadata - await asyncio.sleep(0.05) - finish_chunk = copy.deepcopy(chunk_template) - finish_chunk["choices"][0]["delta"] = {} - finish_chunk["choices"][0]["finish_reason"] = last_finish_reason or "stop" - finish_chunk["extra"] = filtered_json - yield f"data: {json.dumps(finish_chunk)}\n\n" - - yield "data: [DONE]\n\n" - continue - - data = json.loads(line[len("data: ") :]) - data["model"] = model_name - - choice = data.get("choices", [{}])[0] - delta = choice.get("delta", {}) - content = delta.get("content", "") or "" - finish_reason = choice.get("finish_reason") - - if finish_reason: - # Save finish_reason, don't forward — we emit it at the end - last_finish_reason = finish_reason - chunk_template = data - elif content: - chunk_template = data - pending += content - - if len(pending) <= buffer_size: - continue - - # Strip tags from the whole pending; emit the prefix that lies safely - # outside the look-ahead window. Tags inside the window stay buffered - # until they're either confirmed (anchored by \n) or completed at DONE. - cleaned, _, _ = _strip_sources_tags(pending) - safe_end = max(0, len(cleaned) - buffer_size) - if safe_end > emitted_len: - # Shallow rebuild: data is fresh from json.loads (no aliasing), - # so we only need to avoid mutating shared inner dicts. - choice = data["choices"][0] - out = { - **data, - "choices": [ - {**choice, "delta": {**choice.get("delta", {}), "content": cleaned[emitted_len:safe_end]}} - ], - "extra": "{}", - } - yield f"data: {json.dumps(out)}\n\n" - emitted_len = safe_end - else: - # Forward non-content, non-finish chunks immediately (e.g. role delta) - data["extra"] = "{}" - yield f"data: {json.dumps(data)}\n\n" - - -# Initialize language detector -lang_detect_cache_dir = "/app/model_weights/" -lang_detector_config = LangDetectConfig( - max_input_length=1024, # chars - model="auto", - cache_dir=lang_detect_cache_dir, -) -lang_detector: LangDetector = LangDetector(config=lang_detector_config) - - -def detect_language(text: str): - outputs = lang_detector.detect(text, k=1) - return outputs[0].get("lang") - - -def get_llm_semaphore() -> DistributedSemaphore: - config = load_config() - return DistributedSemaphore( - name="llmSemaphore", - max_concurrent_ops=config.semaphore.llm_semaphore, - ) - - -def get_vlm_semaphore() -> DistributedSemaphore: - config = load_config() - return DistributedSemaphore( - name="vlmSemaphore", - max_concurrent_ops=config.semaphore.vlm_semaphore, - ) - - -def get_audio_semaphore() -> DistributedSemaphore: - config = load_config() - return DistributedSemaphore( - name="audioSemaphore", - max_concurrent_ops=config.loader.transcriber.max_concurrent_chunks, - ) diff --git a/openrag/config/__init__.py b/openrag/config/__init__.py deleted file mode 100644 index b7b6b5790..000000000 --- a/openrag/config/__init__.py +++ /dev/null @@ -1,39 +0,0 @@ -# Re-export from canonical location for backward compatibility. -# New code should import from openrag.core.config directly. -"""OpenRAG configuration package. - -Public API: - load_config() — load config (cached singleton, or fresh with overrides) - Settings — root Pydantic model - get_settings() — cached singleton accessor -""" - -from functools import lru_cache - -from openrag.core.config.root import Settings # noqa: F401 - - -@lru_cache -def get_settings() -> Settings: - """Cached singleton — one Settings instance per process.""" - from openrag.core.config.loader import load_config as _load - - return _load() - - -def load_config(config_path=None, overrides=None) -> Settings: - """Return the cached Pydantic Settings singleton. - - The ``config_path`` parameter is kept for backward compatibility. - Use ``OPENRAG_CONF_DIR`` env var to override the config directory. - - The ``overrides`` parameter bypasses the cache (useful for tests). - """ - if overrides or config_path: - from openrag.core.config.loader import load_config as _load - - return _load(conf_dir=config_path, overrides=overrides) - return get_settings() - - -__all__ = ["load_config", "Settings", "get_settings"] diff --git a/openrag/config/loader.py b/openrag/config/loader.py deleted file mode 100644 index a8c5007d5..000000000 --- a/openrag/config/loader.py +++ /dev/null @@ -1,296 +0,0 @@ -"""Configuration loader — reads YAML defaults, merges env var overrides, validates with Pydantic.""" - -from __future__ import annotations - -import logging -import os -from pathlib import Path -from typing import Any - -import yaml - -from openrag.core.config.root import Settings - -logger = logging.getLogger(__name__) - -_DEFAULT_CONF_DIR = Path(__file__).resolve().parent.parent.parent / "conf" - -# --------------------------------------------------------------------------- -# Env var mappings: {env_var_name: dotted.config.path} -# -# Only values that should be overridable at deploy time are listed here. -# Secrets (API keys, passwords) and deployment-specific values (hosts, ports). -# Operational knobs that ops teams commonly tune are also included. -# --------------------------------------------------------------------------- -_ENV_OVERRIDES: list[tuple[str, str, type]] = [ - # LLM - ("BASE_URL", "llm.base_url", str), - ("MODEL", "llm.model", str), - ("API_KEY", "llm.api_key", str), - # VLM - ("VLM_BASE_URL", "vlm.base_url", str), - ("VLM_MODEL", "vlm.model", str), - ("VLM_API_KEY", "vlm.api_key", str), - # Semaphore - ("LLM_SEMAPHORE", "semaphore.llm_semaphore", int), - ("VLM_SEMAPHORE", "semaphore.vlm_semaphore", int), - # Embedder - ("EMBEDDER_MODEL_NAME", "embedder.model_name", str), - ("EMBEDDER_BASE_URL", "embedder.base_url", str), - ("EMBEDDER_API_KEY", "embedder.api_key", str), - ("MAX_MODEL_LEN", "embedder.max_model_len", int), - # VectorDB - ("VDB_HOST", "vectordb.host", str), - ("VDB_iPORT", "vectordb.port", int), # legacy typo, kept for backward compat - ("VDB_PORT", "vectordb.port", int), # canonical name, wins if both are set - ("VDB_CONNECTOR_NAME", "vectordb.connector_name", str), - ("VDB_COLLECTION_NAME", "vectordb.collection_name", str), - ("VDB_HYBRID_SEARCH", "vectordb.hybrid_search", bool), - ("VDB_ENABLE_INSERTION", "vectordb.enable", bool), - # RDB (Postgres) - ("POSTGRES_HOST", "rdb.host", str), - ("POSTGRES_PORT", "rdb.port", int), - ("POSTGRES_USER", "rdb.user", str), - ("POSTGRES_PASSWORD", "rdb.password", str), - ("DEFAULT_FILE_QUOTA", "rdb.default_file_quota", int), - # Reranker - ("RERANKER_PROVIDER", "reranker.provider", str), - ("RERANKER_ENABLED", "reranker.enabled", bool), - ("RERANKER_MODEL", "reranker.model_name", str), - ("RERANKER_TOP_K", "reranker.top_k", int), - ("RERANKER_BASE_URL", "reranker.base_url", str), - ("RERANKER_API_KEY", "reranker.api_key", str), - ("RERANKER_TIMEOUT", "reranker.timeout", float), - ("RERANKER_SEMAPHORE", "reranker.semaphore", int), - # Map-Reduce - ("MAP_REDUCE_INITIAL_BATCH_SIZE", "map_reduce.initial_batch_size", int), - ("MAP_REDUCE_EXPANSION_BATCH_SIZE", "map_reduce.expansion_batch_size", int), - ("MAP_REDUCE_MAX_TOTAL_DOCUMENTS", "map_reduce.max_total_documents", int), - ("MAP_REDUCE_DEBUG", "map_reduce.debug", bool), - # Verbose - ("LOG_LEVEL", "verbose.level", str), - # Server - ("PREFERRED_URL_SCHEME", "server.preferred_url_scheme", str), - # LLM Context - ("MAX_LLM_CONTEXT_SIZE", "llm_context.max_llm_context_size", int), - ("MAX_OUTPUT_TOKENS", "llm_context.max_output_tokens", int), - # Paths - ("PROMPTS_DIR", "paths.prompts_dir", str), - ("DATA_DIR", "paths.data_dir", str), - ("DB_DIR", "paths.db_dir", str), - ("LOG_DIR", "paths.log_dir", str), - # Loader - ("IMAGE_CAPTIONING", "loader.image_captioning", bool), - ("IMAGE_CAPTIONING_URL", "loader.image_captioning_url", bool), - ("SAVE_MARKDOWN", "loader.save_markdown", bool), - ("PDFLoader", "loader.file_loaders.pdf", str), - ("AUDIOLOADER", "loader.file_loaders.wav", str), - ("MARKER_MAX_TASKS_PER_CHILD", "loader.marker_max_tasks_per_child", int), - ("MARKER_POOL_SIZE", "loader.marker_pool_size", int), - ("MARKER_MAX_PROCESSES", "loader.marker_max_processes", int), - ("MARKER_NUM_GPUS", "loader.marker_num_gpus", float), - ("MARKER_TIMEOUT", "loader.marker_timeout", int), - ("MARKER_PDFTEXT_WORKERS", "loader.marker_pdftext_workers", int), - ("MARKER_CHUNK_SIZE", "loader.marker_chunk_size", int), - ("DOCLING_NUM_GPUS", "loader.docling_num_gpus", float), - ("DOCLING_POOL_SIZE", "loader.docling_pool_size", int), - ("DOCLING_MAX_TASKS_PER_WORKER", "loader.docling_max_tasks_per_worker", int), - ("WHISPER_MODEL", "loader.local_whisper.model", str), - ("WHISPER_N_WORKERS", "loader.local_whisper.whisper_n_workers", int), - ("WHISPER_NUM_GPUS", "loader.local_whisper.whisper_num_gpus", float), - ("WHISPER_CONCURRENCY_PER_WORKER", "loader.local_whisper.whisper_concurrency_per_worker", int), - ("TRANSCRIBER_BASE_URL", "loader.transcriber.base_url", str), - ("TRANSCRIBER_API_KEY", "loader.transcriber.api_key", str), - ("TRANSCRIBER_MODEL", "loader.transcriber.model_name", str), - ("TRANSCRIBER_TIMEOUT", "loader.transcriber.timeout", int), - ("TRANSCRIBER_MAX_CONCURRENT_CHUNKS", "loader.transcriber.max_concurrent_chunks", int), - ("USE_WHISPER_LANG_DETECTOR", "loader.transcriber.use_whisper_lang_detector", bool), - ("TRANSCRIBER_DIRECT_UPLOAD_SUFFIXES", "loader.transcriber.direct_upload_suffixes", str), - ("OPENAI_LOADER_BASE_URL", "loader.openai.base_url", str), - ("OPENAI_LOADER_API_KEY", "loader.openai.api_key", str), - ("OPENAI_LOADER_MODEL", "loader.openai.model", str), - ("OPENAI_LOADER_TEMPERATURE", "loader.openai.temperature", float), - ("OPENAI_LOADER_TIMEOUT", "loader.openai.timeout", int), - ("OPENAI_LOADER_MAX_RETRIES", "loader.openai.max_retries", int), - ("OPENAI_LOADER_TOP_P", "loader.openai.top_p", float), - ("OPENAI_LOADER_CONCURRENCY_LIMIT", "loader.openai.concurrency_limit", int), - # Ray - ("RAY_NUM_GPUS", "ray.num_gpus", float), - ("RAY_POOL_SIZE", "ray.pool_size", int), - ("RAY_MAX_TASKS_PER_WORKER", "ray.max_tasks_per_worker", int), - ("RAY_MAX_TASK_RETRIES", "ray.indexer.max_task_retries", int), - ("INDEXER_SERIALIZE_TIMEOUT", "ray.indexer.serialize_timeout", int), - ("VECTORDB_TIMEOUT", "ray.indexer.vectordb_timeout", int), - ("INDEXER_DEFAULT_CONCURRENCY", "ray.indexer.concurrency_groups.default", int), - ("INDEXER_UPDATE_CONCURRENCY", "ray.indexer.concurrency_groups.update", int), - ("INDEXER_SEARCH_CONCURRENCY", "ray.indexer.concurrency_groups.search", int), - ("INDEXER_DELETE_CONCURRENCY", "ray.indexer.concurrency_groups.delete", int), - ("INDEXER_SERIALIZE_CONCURRENCY", "ray.indexer.concurrency_groups.serialize", int), - ("INDEXER_CHUNK_CONCURRENCY", "ray.indexer.concurrency_groups.chunk", int), - ("INDEXER_INSERT_CONCURRENCY", "ray.indexer.concurrency_groups.insert", int), - ("RAY_SEMAPHORE_CONCURRENCY", "ray.semaphore.concurrency", int), - ("ENABLE_RAY_SERVE", "ray.serve.enable", bool), - ("RAY_SERVE_NUM_REPLICAS", "ray.serve.num_replicas", int), - ("RAY_SERVE_HOST", "ray.serve.host", str), - ("RAY_SERVE_PORT", "ray.serve.port", int), - ("CHAINLIT_PORT", "ray.serve.chainlit_port", int), - # Chunker - ("CHUNKER", "chunker.name", str), - ("CONTEXTUAL_RETRIEVAL", "chunker.contextual_retrieval", bool), - ("CONTEXTUALIZATION_TIMEOUT", "chunker.contextualization_timeout", int), - ("MAX_CONCURRENT_CONTEXTUALIZATION", "chunker.max_concurrent_contextualization", int), - ("CHUNK_SIZE", "chunker.chunk_size", int), - ("CHUNK_OVERLAP_RATE", "chunker.chunk_overlap_rate", float), - # Retriever - ("RETRIEVER_TYPE", "retriever.type", str), - ("RETRIEVER_TOP_K", "retriever.top_k", int), - ("SIMILARITY_THRESHOLD", "retriever.similarity_threshold", float), - ("WITH_SURROUNDING_CHUNKS", "retriever.with_surrounding_chunks", bool), - ("INCLUDE_RELATED", "retriever.include_related", bool), - ("INCLUDE_ANCESTORS", "retriever.include_ancestors", bool), - ("RELATED_LIMIT", "retriever.related_limit", int), - ("MAX_DEPTH", "retriever.max_ancestor_depth", int), - ("RETRIEVER_ALLOW_FILTERLESS_FALLBACK", "retriever.allow_filterless_fallback", bool), - # RAG - ("RAG_MODE", "rag.mode", str), - # WebSearch - ("WEBSEARCH_PROVIDER", "websearch.provider", str), - ("WEBSEARCH_API_TOKEN", "websearch.api_token", str), - ("WEBSEARCH_BASE_URL", "websearch.base_url", str), - ("WEBSEARCH_TOP_K", "websearch.top_k", int), - ("WEBSEARCH_LANG", "websearch.lang", str), - ("WEBSEARCH_MAX_TOKENS", "websearch.max_tokens", int), - ("WEBSEARCH_FETCH_CONTENT", "websearch.fetch_content", bool), - ("WEBSEARCH_FETCH_MAX_RESULTS", "websearch.fetch_max_results", int), - ("WEBSEARCH_FETCH_TIMEOUT", "websearch.fetch_timeout", float), - ("WEBSEARCH_FETCH_MAX_TOKENS", "websearch.fetch_max_tokens", int), - ("WEBSEARCH_FETCH_VERIFY_SSL", "websearch.fetch_verify_ssl", bool), -] - -# Audio loader env var applies to all audio/video extensions -_AUDIO_EXTENSIONS = ("mp3", "flac", "ogg", "aac", "flv", "wma", "mp4") - - -def _load_yaml(path: Path) -> dict[str, Any]: - """Load a YAML file, returning empty dict if not found.""" - if not path.exists(): - logger.warning("Config file not found: %s — using defaults", path) - return {} - with open(path) as f: - data = yaml.safe_load(f) - return data or {} - - -def _deep_merge(base: dict, override: dict) -> dict: - """Recursively merge override into base.""" - merged = base.copy() - for key, value in override.items(): - if key in merged and isinstance(merged[key], dict) and isinstance(value, dict): - merged[key] = _deep_merge(merged[key], value) - else: - merged[key] = value - return merged - - -def _set_nested(data: dict, dotted_path: str, value: Any) -> None: - """Set a value in a nested dict using a dotted path like 'ray.indexer.timeout'.""" - keys = dotted_path.split(".") - current = data - for key in keys[:-1]: - current = current.setdefault(key, {}) - current[keys[-1]] = value - - -def _coerce(value: str, target_type: type, env_var: str = "") -> Any: - """Coerce a string env var value to the target type.""" - if target_type is bool: - lower = value.lower() - if lower in ("true", "1", "yes"): - return True - if lower in ("false", "0", "no"): - return False - raise ValueError(f"Invalid value for {env_var}: expected bool, got {value!r}") - try: - if target_type is int: - return int(value) - if target_type is float: - return float(value) - except ValueError: - raise ValueError(f"Invalid value for {env_var}: expected {target_type.__name__}, got {value!r}") - return value - - -def _apply_env_overrides(data: dict) -> dict: - """Apply environment variable overrides to the config dict.""" - for env_var, dotted_path, target_type in _ENV_OVERRIDES: - value = os.environ.get(env_var) - if value is not None and value != "": - _set_nested(data, dotted_path, _coerce(value, target_type, env_var)) - - # SEMAPHORE sets both LLM and VLM semaphores (convenience shorthand) - semaphore = os.environ.get("SEMAPHORE") - if semaphore: - sem_value = _coerce(semaphore, int, "SEMAPHORE") - sem = data.setdefault("semaphore", {}) - sem.setdefault("llm_semaphore", sem_value) - sem.setdefault("vlm_semaphore", sem_value) - - # AUDIOLOADER applies to all audio/video extensions - audio_loader = os.environ.get("AUDIOLOADER") - if audio_loader: - file_loaders = data.setdefault("loader", {}).setdefault("file_loaders", {}) - for ext in _AUDIO_EXTENSIONS: - file_loaders[ext] = audio_loader - - return data - - -def load_config( - conf_dir: Path | str | None = None, - overrides: dict[str, Any] | None = None, -) -> Settings: - """Load configuration: YAML defaults → env var overrides → Pydantic validation. - - Args: - conf_dir: Path to the configuration directory. Defaults to ``conf/`` - at the project root, overridable via ``OPENRAG_CONF_DIR``. - overrides: Programmatic overrides (useful for tests). - """ - from dotenv import load_dotenv - - load_dotenv() - - env_conf_dir = os.environ.get("OPENRAG_CONF_DIR") - if conf_dir: - conf_dir = Path(conf_dir) - elif env_conf_dir: - conf_dir = Path(env_conf_dir) - else: - conf_dir = _DEFAULT_CONF_DIR - - # 1. Load YAML defaults - data = _load_yaml(conf_dir / "config.yaml") - - # Remove YAML anchors (keys starting with _) — they are DRY helpers, not config sections - data = {k: v for k, v in data.items() if not k.startswith("_")} - - # 2. Apply env var overrides - data = _apply_env_overrides(data) - - # Strip blank reranker.base_url so the provider-specific Pydantic default applies - reranker = data.get("reranker") - if isinstance(reranker, dict) and not reranker.get("base_url"): - reranker.pop("base_url", None) - - # 3. Apply programmatic overrides (tests) - if overrides: - data = _deep_merge(data, overrides) - - # 4. Resolve paths (after all merging so overrides are honored) - paths = data.get("paths", {}) - for key in ("prompts_dir", "data_dir", "db_dir", "log_dir"): - if key in paths and paths[key]: - paths[key] = str(Path(paths[key]).resolve()) - - # 5. Validate with Pydantic - return Settings(**data) diff --git a/openrag/config/models.py b/openrag/config/models.py deleted file mode 100644 index 46600c246..000000000 --- a/openrag/config/models.py +++ /dev/null @@ -1,538 +0,0 @@ -"""Pydantic config models — pure validation schemas. - -Each model corresponds to a configuration section. Defaults are fallbacks only; -in production, values come from conf/config.yaml merged with env var overrides -(see loader.py for the merge logic). -""" - -from __future__ import annotations - -from pathlib import Path -from typing import Annotated, Any, Literal - -from pydantic import BaseModel, Field, field_validator - - -# --------------------------------------------------------------------------- -# Base mixin — frozen models with dict-like backward compat -# --------------------------------------------------------------------------- -class ConfigMixin(BaseModel): - """Frozen Pydantic model with dict-like access for backward compatibility. - - Existing code using ``config.section.get("key")``, ``config.section["key"]``, - ``dict(config.section)``, and ``**config.section`` keeps working. - """ - - model_config = {"frozen": True} - - def get(self, key: str, default: Any = None) -> Any: - try: - return getattr(self, key) - except AttributeError: - return default - - def __getitem__(self, key: str) -> Any: - try: - return getattr(self, key) - except AttributeError: - raise KeyError(key) - - def keys(self): - return list(type(self).model_fields.keys()) - - def values(self): - return [getattr(self, k) for k in type(self).model_fields] - - def items(self): - return [(k, getattr(self, k)) for k in type(self).model_fields] - - def __iter__(self): - return iter(type(self).model_fields) - - def __contains__(self, key: str) -> bool: - return key in type(self).model_fields - - -# --------------------------------------------------------------------------- -# LLM params (shared by llm and vlm) -# --------------------------------------------------------------------------- -class LLMParamsConfig(ConfigMixin): - temperature: float = 0.1 - timeout: int = 60 - max_retries: int = 2 - logprobs: bool = True - - -# --------------------------------------------------------------------------- -# LLM -# --------------------------------------------------------------------------- -class LLMConfig(LLMParamsConfig): - base_url: str = "" - model: str = "" - api_key: str = Field(default="", repr=False) - - -# --------------------------------------------------------------------------- -# VLM -# --------------------------------------------------------------------------- -class VLMConfig(LLMParamsConfig): - base_url: str = "" - model: str = "" - api_key: str = Field(default="", repr=False) - - -# --------------------------------------------------------------------------- -# Semaphore -# --------------------------------------------------------------------------- -class SemaphoreConfig(ConfigMixin): - llm_semaphore: int = 10 - vlm_semaphore: int = 10 - - -# --------------------------------------------------------------------------- -# Embedder -# --------------------------------------------------------------------------- -class EmbedderConfig(ConfigMixin): - provider: str = "openai" - model_name: str = "jinaai/jina-embeddings-v3" - base_url: str = "http://vllm:8000/v1" - api_key: str = Field(default="EMPTY", repr=False) - max_model_len: int = 8192 - - -# --------------------------------------------------------------------------- -# VectorDB -# --------------------------------------------------------------------------- -class VectorDBConfig(ConfigMixin): - host: str = "milvus" - port: int = 19530 - connector_name: str = "milvus" - collection_name: str = "vdb_test" - hybrid_search: bool = True - enable: bool = True - schema_version: int = 1 - - -# --------------------------------------------------------------------------- -# RDB (Postgres) -# --------------------------------------------------------------------------- -class RDBConfig(ConfigMixin): - host: str = "rdb" - port: int = 5432 - user: str = "root" - password: str = Field(default="", repr=False) - default_file_quota: int = -1 - - -# --------------------------------------------------------------------------- -# Reranker -# --------------------------------------------------------------------------- -class _BaseRerankerConfig(ConfigMixin): - model_name: str = "Alibaba-NLP/gte-multilingual-reranker-base" - top_k: int = 10 - api_key: str = Field(default="EMPTY", repr=False) - timeout: float = 60.0 - semaphore: int = 5 - enabled: bool = True - - -class InfinityRerankerConfig(_BaseRerankerConfig): - provider: Literal["infinity"] = "infinity" - base_url: str = "http://reranker:7997" - - -class OpenAIRerankerConfig(_BaseRerankerConfig): - provider: Literal["openai"] = "openai" - base_url: str = "http://reranker:8000/v1" - - -RerankerConfig = Annotated[ - InfinityRerankerConfig | OpenAIRerankerConfig, - Field(discriminator="provider"), -] - - -def _default_reranker_config() -> InfinityRerankerConfig: - return InfinityRerankerConfig() - - -# --------------------------------------------------------------------------- -# MapReduce -# --------------------------------------------------------------------------- -class MapReduceConfig(ConfigMixin): - initial_batch_size: int = 10 - expansion_batch_size: int = 5 - max_total_documents: int = 20 - debug: bool = False - - -# --------------------------------------------------------------------------- -# Verbose -# --------------------------------------------------------------------------- -class VerboseConfig(ConfigMixin): - level: str = "DEBUG" - - -# --------------------------------------------------------------------------- -# Server -# --------------------------------------------------------------------------- -class ServerConfig(ConfigMixin): - preferred_url_scheme: str | None = None - - -# --------------------------------------------------------------------------- -# LLM Context -# --------------------------------------------------------------------------- -class LLMContextConfig(ConfigMixin): - max_llm_context_size: int = 8192 - max_output_tokens: int = 1024 - - -# --------------------------------------------------------------------------- -# Paths -# --------------------------------------------------------------------------- -class PathsConfig(ConfigMixin): - prompts_dir: Path = Path("../prompts/example1") - data_dir: Path = Path("../data") - db_dir: Path = Path("/app/db") - log_dir: Path = Path("/app/logs") - - model_config = {**ConfigMixin.model_config, "arbitrary_types_allowed": True} - - -# --------------------------------------------------------------------------- -# Prompts -# --------------------------------------------------------------------------- -class PromptsConfig(ConfigMixin): - sys_prompt: str = "sys_prompt_tmpl.txt" - query_contextualizer: str = "query_contextualizer_tmpl.txt" - chunk_contextualizer: str = "chunk_contextualizer_tmpl.txt" - image_describer: str = "image_captioning_tmpl.txt" - spoken_style_answer: str = "spoken_style_answer_tmpl.txt" - hyde: str = "hyde.txt" - multi_query: str = "multi_query_pmpt_tmpl.txt" - - -# --------------------------------------------------------------------------- -# Transcriber (nested under loader) -# --------------------------------------------------------------------------- -_DEFAULT_DIRECT_UPLOAD_SUFFIXES = frozenset( - {".wav", ".flac", ".ogg", ".mp3", ".mp4", ".m4a", ".webm", ".mpeg", ".mpga"} -) - - -def _normalize_suffix(s: str) -> str: - s = s.strip().lower() - if not s: - return "" - return s if s.startswith(".") else f".{s}" - - -class TranscriberConfig(ConfigMixin): - base_url: str = "http://transcriber:8000/v1" - api_key: str = Field(default="EMPTY", repr=False) - model_name: str = "openai/whisper-large-v3-turbo" - timeout: int = 3600 - max_concurrent_chunks: int = 20 - use_whisper_lang_detector: bool = True - direct_upload_suffixes: set[str] = Field(default_factory=lambda: set(_DEFAULT_DIRECT_UPLOAD_SUFFIXES)) - - @field_validator("direct_upload_suffixes", mode="before") - @classmethod - def _split_suffixes(cls, v: Any) -> Any: - if isinstance(v, str): - return {n for raw in v.split("|") if (n := _normalize_suffix(raw))} - return v - - -# --------------------------------------------------------------------------- -# OpenAI Loader (nested under loader) -# --------------------------------------------------------------------------- -class OpenAILoaderConfig(ConfigMixin): - base_url: str = "http://openai:8000/v1" - api_key: str = Field(default="EMPTY", repr=False) - model: str = "dotsocr-model" - temperature: float = 0.2 - timeout: int = 180 - max_retries: int = 2 - top_p: float = 0.9 - concurrency_limit: int = 20 - - -# --------------------------------------------------------------------------- -# Local Whisper (nested under loader) -# --------------------------------------------------------------------------- -class LocalWhisperConfig(ConfigMixin): - model: str = "base" - whisper_n_workers: int = 3 - whisper_num_gpus: float = 0.01 - whisper_concurrency_per_worker: int = 2 - whisper_timeout: int = 1800 - whisper_max_task_retry: int = 1 - whisper_retry_base_delay: float = 2.0 - - -# --------------------------------------------------------------------------- -# File loaders mapping (nested under loader) -# --------------------------------------------------------------------------- -class FileLoadersConfig(ConfigMixin): - txt: str = "TextLoader" - pdf: str = "MarkerLoader" - eml: str = "EmlLoader" - docx: str = "DocxLoader" - pptx: str = "PPTXLoader" - doc: str = "DocLoader" - png: str = "ImageLoader" - jpeg: str = "ImageLoader" - jpg: str = "ImageLoader" - svg: str = "ImageLoader" - wav: str = "LocalWhisperLoader" - mp3: str = "LocalWhisperLoader" - flac: str = "LocalWhisperLoader" - ogg: str = "LocalWhisperLoader" - aac: str = "LocalWhisperLoader" - flv: str = "LocalWhisperLoader" - wma: str = "LocalWhisperLoader" - mp4: str = "LocalWhisperLoader" - md: str = "MarkdownLoader" - - -# --------------------------------------------------------------------------- -# Mimetypes mapping (nested under loader) -# --------------------------------------------------------------------------- -class MimetypesConfig(ConfigMixin): - """Maps MIME type strings to file extensions. - - Stored as regular fields so Pydantic serialization works normally. - Access via .to_dict() for {mime_type: extension} mapping. - """ - - text_plain: str = Field(default=".txt", alias="text/plain") - text_markdown: str = Field(default=".md", alias="text/markdown") - application_pdf: str = Field(default=".pdf", alias="application/pdf") - message_rfc822: str = Field(default=".eml", alias="message/rfc822") - application_docx: str = Field( - default=".docx", - alias="application/vnd.openxmlformats-officedocument.wordprocessingml.document", - ) - application_pptx: str = Field( - default=".pptx", - alias="application/vnd.openxmlformats-officedocument.presentationml.presentation", - ) - application_msword: str = Field(default=".doc", alias="application/msword") - image_png: str = Field(default=".png", alias="image/png") - image_jpeg: str = Field(default=".jpeg", alias="image/jpeg") - audio_wav: str = Field(default=".wav", alias="audio/wav") - audio_mpeg: str = Field(default=".mp3", alias="audio/mpeg") - audio_flac: str = Field(default=".flac", alias="audio/flac") - audio_ogg: str = Field(default=".ogg", alias="audio/ogg") - audio_aac: str = Field(default=".aac", alias="audio/aac") - video_x_flv: str = Field(default=".flv", alias="video/x-flv") - audio_x_ms_wma: str = Field(default=".wma", alias="audio/x-ms-wma") - video_mp4: str = Field(default=".mp4", alias="video/mp4") - - model_config = {"frozen": True, "extra": "allow", "populate_by_name": True} - - def to_dict(self) -> dict[str, str]: - """Return {mime_type: extension} mapping using aliases as keys.""" - result = {} - for field_name, field_info in type(self).model_fields.items(): - alias = field_info.alias or field_name - result[alias] = getattr(self, field_name) - if self.__pydantic_extra__: - result.update(self.__pydantic_extra__) - return result - - -# --------------------------------------------------------------------------- -# Loader -# --------------------------------------------------------------------------- -class LoaderConfig(ConfigMixin): - image_captioning: bool = True - image_captioning_url: bool = True - save_markdown: bool = False - mimetypes: MimetypesConfig = Field(default_factory=MimetypesConfig) - local_whisper: LocalWhisperConfig = Field(default_factory=LocalWhisperConfig) - file_loaders: FileLoadersConfig = Field(default_factory=FileLoadersConfig) - marker_max_tasks_per_child: int = 20 - marker_pool_size: int = 1 - marker_max_processes: int = 2 - marker_num_gpus: float = 0.01 - marker_timeout: int = 3600 - marker_pdftext_workers: int = 2 - marker_chunk_size: int = 10 - marker_max_task_retry: int = 3 - marker_retry_base_delay: float = 2.0 - docling_num_gpus: float = Field(default=0.01, ge=0) - docling_pool_size: int = Field(default=1, ge=1) - docling_max_tasks_per_worker: int = Field(default=2, ge=1) - docling_timeout: int = 3600 - docling_max_task_retry: int = 3 - docling_retry_base_delay: float = 2.0 - transcriber: TranscriberConfig = Field(default_factory=TranscriberConfig) - openai: OpenAILoaderConfig = Field(default_factory=OpenAILoaderConfig) - # Max depth of nested .eml-in-.eml attachments the EmlLoader will descend - # into. Bounds recursion when .eml files are nested inside one another. - eml_max_recursion_depth: int = 5 - - -# --------------------------------------------------------------------------- -# Ray — Indexer concurrency groups -# --------------------------------------------------------------------------- -class IndexerConcurrencyGroupsConfig(ConfigMixin): - default: int = 1000 - update: int = 100 - search: int = 100 - delete: int = 100 - serialize: int = 50 - chunk: int = 1000 - insert: int = 100 - - -class RayIndexerConfig(ConfigMixin): - max_task_retries: int = 2 - serialize_timeout: int = 3600 - vectordb_timeout: int = 30 - concurrency_groups: IndexerConcurrencyGroupsConfig = Field( - default_factory=IndexerConcurrencyGroupsConfig, - ) - - -class RaySemaphoreConfig(ConfigMixin): - concurrency: int = 100000 - - -class RayServeConfig(ConfigMixin): - enable: bool = False - num_replicas: int = 1 - host: str = "0.0.0.0" - port: int = 8080 - chainlit_port: int = 8090 - - -class RayConfig(ConfigMixin): - num_gpus: float = 0.01 - pool_size: int = 1 - max_tasks_per_worker: int = 8 - indexer: RayIndexerConfig = Field(default_factory=RayIndexerConfig) - semaphore: RaySemaphoreConfig = Field(default_factory=RaySemaphoreConfig) - serve: RayServeConfig = Field(default_factory=RayServeConfig) - - -# --------------------------------------------------------------------------- -# Chunker -# --------------------------------------------------------------------------- -class ChunkerConfig(ConfigMixin): - name: str = "recursive_splitter" - contextual_retrieval: bool = True - contextualization_timeout: int = 120 - max_concurrent_contextualization: int = 10 - chunk_size: int = 512 - chunk_overlap_rate: float = 0.2 - - -# --------------------------------------------------------------------------- -# Retriever -# --------------------------------------------------------------------------- -class _BaseRetrieverConfig(ConfigMixin): - top_k: int = 50 - similarity_threshold: float = 0.6 - with_surrounding_chunks: bool = False - include_related: bool = True - include_ancestors: bool = True - related_limit: int = 10 - max_ancestor_depth: int = 10 - # Absolute upper bound on ancestor-CTE recursion regardless of caller- - # supplied depth — defends against cyclic parent_id chains that would - # otherwise loop until the DB aborts the query. Operators can tune via - # the retriever config; the default is well above realistic deep - # document hierarchies. - max_ancestor_depth_cap: int = 1000 - allow_filterless_fallback: bool = True - - -class SingleRetrieverConfig(_BaseRetrieverConfig): - type: Literal["single"] = "single" - - -class MultiQueryRetrieverConfig(_BaseRetrieverConfig): - type: Literal["multiQuery"] = "multiQuery" - k_queries: int = 3 - - -class HydeRetrieverConfig(_BaseRetrieverConfig): - type: Literal["hyde"] = "hyde" - combine: bool = False - - -RetrieverConfig = Annotated[ - SingleRetrieverConfig | MultiQueryRetrieverConfig | HydeRetrieverConfig, - Field(discriminator="type"), -] - - -# --------------------------------------------------------------------------- -# RAG -# --------------------------------------------------------------------------- -class RAGConfig(ConfigMixin): - mode: str = "ChatBotRag" - chat_history_depth: int = 4 - max_contextualized_query_len: int = 512 - - -# --------------------------------------------------------------------------- -# WebSearch -# --------------------------------------------------------------------------- -class _BaseWebSearchConfig(ConfigMixin): - base_url: str - api_token: str = Field(default="", repr=False) - top_k: int = 5 - lang: str = "fr-FR" - max_tokens: int = 2000 - fetch_content: bool = True - fetch_max_results: int = 3 - fetch_timeout: float = 1.0 - fetch_max_tokens: int = 500 - fetch_verify_ssl: bool = True - - -class StaanWebSearchConfig(_BaseWebSearchConfig): - provider: Literal["staan"] = "staan" - base_url: str = "https://api.staan.ai/search/web" - - -WebSearchConfig = Annotated[ - StaanWebSearchConfig, - Field(discriminator="provider"), -] - - -# --------------------------------------------------------------------------- -# Root Settings — composes all sub-models -# --------------------------------------------------------------------------- -class Settings(ConfigMixin): - """Root configuration. - - Defaults here are fallbacks only. In production, values come from - conf/config.yaml merged with environment variable overrides. - """ - - llm: LLMConfig = Field(default_factory=LLMConfig) - vlm: VLMConfig = Field(default_factory=VLMConfig) - semaphore: SemaphoreConfig = Field(default_factory=SemaphoreConfig) - embedder: EmbedderConfig = Field(default_factory=EmbedderConfig) - vectordb: VectorDBConfig = Field(default_factory=VectorDBConfig) - rdb: RDBConfig = Field(default_factory=RDBConfig) - reranker: RerankerConfig = Field(default_factory=_default_reranker_config) - map_reduce: MapReduceConfig = Field(default_factory=MapReduceConfig) - verbose: VerboseConfig = Field(default_factory=VerboseConfig) - server: ServerConfig = Field(default_factory=ServerConfig) - llm_context: LLMContextConfig = Field(default_factory=LLMContextConfig) - paths: PathsConfig = Field(default_factory=PathsConfig) - prompts: PromptsConfig = Field(default_factory=PromptsConfig) - loader: LoaderConfig = Field(default_factory=LoaderConfig) - ray: RayConfig = Field(default_factory=RayConfig) - chunker: ChunkerConfig = Field(default_factory=ChunkerConfig) - retriever: RetrieverConfig = Field(default_factory=SingleRetrieverConfig) - rag: RAGConfig = Field(default_factory=RAGConfig) - websearch: WebSearchConfig = Field(default_factory=StaanWebSearchConfig) diff --git a/openrag/core/utils/singleton.py b/openrag/core/utils/singleton.py new file mode 100644 index 000000000..5dd8c43e5 --- /dev/null +++ b/openrag/core/utils/singleton.py @@ -0,0 +1,9 @@ +class SingletonMeta(type): + """Simple singleton metaclass for legacy adapter objects.""" + + _instances = {} + + def __call__(cls, *args, **kwargs): + if cls not in cls._instances: + cls._instances[cls] = super().__call__(*args, **kwargs) + return cls._instances[cls] diff --git a/openrag/core/utils/source_filtering.py b/openrag/core/utils/source_filtering.py index ff3d674fe..960792b7c 100644 --- a/openrag/core/utils/source_filtering.py +++ b/openrag/core/utils/source_filtering.py @@ -11,11 +11,19 @@ logger = get_logger() +_EMAIL_RE = re.compile(r"(? str: + preview = _EMAIL_RE.sub("***@***", text) + if len(preview) > max_length: + preview = preview[-max_length:] + return preview def _strip_sources_tags(text: str) -> tuple[str, set[int], bool]: @@ -35,7 +43,7 @@ def extract_and_strip_sources_block(text: str) -> tuple[str, set[int] | None]: if not citations and not saw_none: tail = text[-150:] if len(text) > 150 else text - logger.debug("No [Sources: ...] tag found in LLM response", tail=repr(tail)) + logger.debug("No [Sources: ...] tag found in LLM response", tail=repr(_sanitize_log_preview(tail))) return text, None cleaned = cleaned.rstrip() diff --git a/openrag/core/utils/test_source_filtering.py b/openrag/core/utils/test_source_filtering.py index 9140886ad..803229103 100644 --- a/openrag/core/utils/test_source_filtering.py +++ b/openrag/core/utils/test_source_filtering.py @@ -5,6 +5,7 @@ import pytest from core.utils.source_filtering import ( _min_sources_tag_buffer_size, + _sanitize_log_preview, extract_and_strip_sources_block, filter_sources_by_citations, stream_with_source_filtering, @@ -101,6 +102,17 @@ def test_sources_none_capitalized(self): assert clean == "Answer text" assert citations == set() + def test_sources_numbers_case_insensitive(self): + text = "Answer text\n[sources: 1, 3]" + clean, citations = extract_and_strip_sources_block(text) + assert clean == "Answer text" + assert citations == {1, 3} + + def test_log_preview_redacts_email_addresses(self): + preview = _sanitize_log_preview("Contact alice@example.com for details") + assert "alice@example.com" not in preview + assert "***@***" in preview + def test_multiple_line_terminal_tags_stripped(self): """Bullet-leak case: LLM emits [Sources: X] per bullet item instead of once at end.""" text = "- Claim one about the codebase.\n[Sources: 1]\n- Claim two about APEX.\n[Sources: 1, 5]\n" diff --git a/openrag/scripts/backup.py b/openrag/scripts/backup.py index 020ec4fd1..3dcc16c34 100644 --- a/openrag/scripts/backup.py +++ b/openrag/scripts/backup.py @@ -212,7 +212,7 @@ def load_openrag_config(logger): Returns: tuple: (RDBConfig, VectorDBConfig) Pydantic config models. """ - from config import load_config + from core.config import load_config try: config = load_config() diff --git a/openrag/scripts/restore.py b/openrag/scripts/restore.py index 75ebcdb32..aa622d725 100644 --- a/openrag/scripts/restore.py +++ b/openrag/scripts/restore.py @@ -261,7 +261,7 @@ def load_openrag_config(logger: Any): Returns: tuple: (RDBConfig, VectorDBConfig) Pydantic config models. """ - from config import load_config + from core.config import load_config try: config = load_config() diff --git a/openrag/services/inference/parsers/test_openai_audio.py b/openrag/services/inference/parsers/test_openai_audio.py index 75c200e51..4389e2f52 100644 --- a/openrag/services/inference/parsers/test_openai_audio.py +++ b/openrag/services/inference/parsers/test_openai_audio.py @@ -21,7 +21,7 @@ if "pydub" not in sys.modules: try: import pydub # noqa: E402,F401 — prefer the real library when it imports cleanly - except Exception: + except (ImportError, ModuleNotFoundError): fake_pydub = types.ModuleType("pydub") fake_pydub.AudioSegment = MagicMock() # type: ignore[attr-defined] sys.modules["pydub"] = fake_pydub diff --git a/openrag/services/inference/runtime.py b/openrag/services/inference/runtime.py index 166ed3b92..d027387af 100644 --- a/openrag/services/inference/runtime.py +++ b/openrag/services/inference/runtime.py @@ -16,8 +16,17 @@ def detect_language(text: str): """Detect the primary language of ``text``.""" - outputs = _lang_detector.detect(text, k=1) - return outputs[0].get("lang") + normalized = text.strip() if isinstance(text, str) else "" + if not normalized: + return None + try: + outputs = _lang_detector.detect(normalized, k=1) + except Exception: + return None + if not outputs: + return None + first = outputs[0] + return first.get("lang") if isinstance(first, dict) else None def get_llm_semaphore() -> DistributedSemaphore: diff --git a/openrag/services/inference/test_runtime.py b/openrag/services/inference/test_runtime.py new file mode 100644 index 000000000..6fd65fa7c --- /dev/null +++ b/openrag/services/inference/test_runtime.py @@ -0,0 +1,34 @@ +from services.inference import runtime + + +class _Detector: + def __init__(self, outputs=None, error: Exception | None = None): + self.outputs = outputs + self.error = error + + def detect(self, text: str, k: int): + if self.error is not None: + raise self.error + return self.outputs + + +def test_detect_language_returns_none_for_blank_input(): + assert runtime.detect_language(" ") is None + + +def test_detect_language_returns_none_when_detector_fails(monkeypatch): + monkeypatch.setattr(runtime, "_lang_detector", _Detector(error=RuntimeError("boom"))) + + assert runtime.detect_language("hello") is None + + +def test_detect_language_returns_none_for_empty_output(monkeypatch): + monkeypatch.setattr(runtime, "_lang_detector", _Detector(outputs=[])) + + assert runtime.detect_language("hello") is None + + +def test_detect_language_returns_lang(monkeypatch): + monkeypatch.setattr(runtime, "_lang_detector", _Detector(outputs=[{"lang": "en"}])) + + assert runtime.detect_language("hello") == "en" diff --git a/openrag/services/persistence/document_repo.py b/openrag/services/persistence/document_repo.py index 421a968e0..5fc487b4f 100644 --- a/openrag/services/persistence/document_repo.py +++ b/openrag/services/persistence/document_repo.py @@ -474,7 +474,7 @@ async def get_file_ancestors( 1000) so a self-referential or cyclic ``parent_id`` chain can't loop indefinitely. """ - from config import load_config + from core.config import load_config hard_cap = int(load_config().retriever.max_ancestor_depth_cap) effective_cap = hard_cap diff --git a/openrag/services/persistence/migrations/alembic/env.py b/openrag/services/persistence/migrations/alembic/env.py index d7954ab5b..d5143b538 100644 --- a/openrag/services/persistence/migrations/alembic/env.py +++ b/openrag/services/persistence/migrations/alembic/env.py @@ -7,7 +7,7 @@ sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) from alembic import context -from config import load_config +from core.config import load_config from services.persistence.schema import metadata as target_metadata from sqlalchemy import URL, engine_from_config, pool diff --git a/openrag/services/persistence/migrations/milvus/1.add_created_at_temporal_fields.py b/openrag/services/persistence/migrations/milvus/1.add_created_at_temporal_fields.py index e0651414c..baef1666a 100644 --- a/openrag/services/persistence/migrations/milvus/1.add_created_at_temporal_fields.py +++ b/openrag/services/persistence/migrations/milvus/1.add_created_at_temporal_fields.py @@ -32,7 +32,7 @@ import argparse import sys -from config import load_config +from core.config import load_config from core.utils.logging import get_logger from pymilvus import DataType, MilvusClient from services.storage.milvus_store import SCHEMA_VERSION_PROPERTY_KEY diff --git a/openrag/services/persistence/migrations/milvus/migrate.py b/openrag/services/persistence/migrations/milvus/migrate.py index 7d3cd8f13..f705aeb44 100644 --- a/openrag/services/persistence/migrations/milvus/migrate.py +++ b/openrag/services/persistence/migrations/milvus/migrate.py @@ -36,7 +36,7 @@ from pathlib import Path from types import ModuleType -from config import load_config +from core.config import load_config from core.utils.logging import get_logger from pymilvus import MilvusClient from services.storage.milvus_store import SCHEMA_VERSION_PROPERTY_KEY diff --git a/openrag/services/persistence/test_ancestor_recursion_cap.py b/openrag/services/persistence/test_ancestor_recursion_cap.py index ff4718f00..3c2ef59cc 100644 --- a/openrag/services/persistence/test_ancestor_recursion_cap.py +++ b/openrag/services/persistence/test_ancestor_recursion_cap.py @@ -13,7 +13,7 @@ async def fetch(self, query: str, *params): def test_hard_cap_lives_in_retriever_config(): - from config import load_config + from core.config import load_config cap = load_config().retriever.max_ancestor_depth_cap assert isinstance(cap, int) @@ -22,7 +22,7 @@ def test_hard_cap_lives_in_retriever_config(): @pytest.mark.asyncio async def test_none_max_depth_is_clamped_to_hard_cap(): - from config import load_config + from core.config import load_config from services.persistence.document_repo import PgDocumentRepository pool = _FakePool() @@ -34,7 +34,7 @@ async def test_none_max_depth_is_clamped_to_hard_cap(): @pytest.mark.asyncio async def test_explicit_depth_above_cap_is_clamped(): - from config import load_config + from core.config import load_config from services.persistence.document_repo import PgDocumentRepository cap = int(load_config().retriever.max_ancestor_depth_cap) diff --git a/openrag/services/workers/indexer_pool.py b/openrag/services/workers/indexer_pool.py index 3b7cafca9..a458c174e 100644 --- a/openrag/services/workers/indexer_pool.py +++ b/openrag/services/workers/indexer_pool.py @@ -13,7 +13,7 @@ class IndexerPool: def __init__(self) -> None: import services.inference.vllm_client # noqa: F401 - from config import load_config + from core.config import load_config from core.embeddings import embedder_registry from services.storage.milvus_store import MilvusVectorStore from services.storage.postgres_store import PostgresStore @@ -102,7 +102,7 @@ def _build_chunker(cfg: Any) -> Any: from core.chunking.factory import create_chunker chunker = create_chunker(cfg) - if not hasattr(chunker, "chunk"): + if not callable(getattr(chunker, "chunk", None)): raise TypeError("Configured chunker does not expose a chunk(document, partition) method") return chunker diff --git a/openrag/services/workers/parsers/doc_serializer.py b/openrag/services/workers/parsers/doc_serializer.py index af53e433d..5fe3e51d9 100644 --- a/openrag/services/workers/parsers/doc_serializer.py +++ b/openrag/services/workers/parsers/doc_serializer.py @@ -18,7 +18,7 @@ @ray.remote(max_restarts=5) class DocSerializer: def __init__(self, data_dir=None, **kwargs) -> None: - from config import load_config + from core.config import load_config from core.utils.logging import get_logger self.logger = get_logger() diff --git a/openrag/services/workers/parsers/doc_serializer_adapter.py b/openrag/services/workers/parsers/doc_serializer_adapter.py index 7842ff38d..525c62dc7 100644 --- a/openrag/services/workers/parsers/doc_serializer_adapter.py +++ b/openrag/services/workers/parsers/doc_serializer_adapter.py @@ -17,7 +17,7 @@ class DocSerializerAdapter(FileSerializer): async def serialize(self, path: str, metadata: dict) -> str: import ray - from config import load_config + from core.config import load_config from services.workers.ray_utils import call_ray_actor_with_timeout cfg = load_config() diff --git a/openrag/services/workers/parsers/docling_workers.py b/openrag/services/workers/parsers/docling_workers.py index 9952777bf..1716516b4 100644 --- a/openrag/services/workers/parsers/docling_workers.py +++ b/openrag/services/workers/parsers/docling_workers.py @@ -14,7 +14,7 @@ import ray import torch -from config import load_config +from core.config import load_config from core.indexing.image_preprocessor import pil_to_png_bytes from core.indexing.parsers.document_parser import BasePooledParser from core.models.document import Document, DocumentType, ImageBlock, ProcessedDocument, TextBlock diff --git a/openrag/services/workers/parsers/legacy_loaders/base.py b/openrag/services/workers/parsers/legacy_loaders/base.py index 8c9018263..e7321567a 100644 --- a/openrag/services/workers/parsers/legacy_loaders/base.py +++ b/openrag/services/workers/parsers/legacy_loaders/base.py @@ -4,8 +4,7 @@ from abc import ABC, abstractmethod from pathlib import Path -from components.prompts import IMAGE_DESCRIBER -from components.utils import get_vlm_semaphore, load_config +from core.config import load_config from core.indexing.image_preprocessor import ( DATA_URI_IMAGE_PATTERN as _CORE_DATA_URI_IMAGE_PATTERN, ) @@ -19,12 +18,14 @@ ensure_png_compatible_mode, # noqa: F401 (re-exported for legacy import path) pil_to_png_bytes, ) +from core.prompts import load_template_by_key from core.utils.external_errors import is_external_resource_error from core.utils.logging import get_logger from langchain_core.messages import HumanMessage from langchain_openai import ChatOpenAI from openai import BadRequestError from PIL import Image +from services.inference.runtime import get_vlm_semaphore from tqdm.asyncio import tqdm logger = get_logger() @@ -52,6 +53,13 @@ def __init__(self, **kwargs) -> None: self.image_captioning = self.config.loader.image_captioning self.image_captioning_url = self.config.loader.image_captioning_url + self.image_describer_prompt = "" + if self.image_captioning or self.image_captioning_url: + self.image_describer_prompt = load_template_by_key( + self.config.paths.prompts_dir, + self.config.prompts, + "image_describer", + ) self.vlm_endpoint = ChatOpenAI(**settings).with_retry(stop_after_attempt=2) @@ -152,7 +160,7 @@ async def get_image_description( "type": "image_url", "image_url": {"url": image_url}, }, - {"type": "text", "text": IMAGE_DESCRIBER}, + {"type": "text", "text": self.image_describer_prompt}, ] ) diff --git a/openrag/services/workers/parsers/legacy_loaders/pdf_loaders/docling.py b/openrag/services/workers/parsers/legacy_loaders/pdf_loaders/docling.py index 5b5d7363d..dd4a7b046 100644 --- a/openrag/services/workers/parsers/legacy_loaders/pdf_loaders/docling.py +++ b/openrag/services/workers/parsers/legacy_loaders/pdf_loaders/docling.py @@ -1,8 +1,8 @@ import asyncio import torch -from components.utils import SingletonMeta from core.utils.logging import get_logger +from core.utils.singleton import SingletonMeta from docling.backend.pypdfium2_backend import PyPdfiumDocumentBackend from docling.datamodel.base_models import InputFormat from docling.datamodel.document import ConversionResult diff --git a/openrag/services/workers/parsers/legacy_loaders/test_doc_loader.py b/openrag/services/workers/parsers/legacy_loaders/test_doc_loader.py index 25e84f50a..aa87a7f0f 100644 --- a/openrag/services/workers/parsers/legacy_loaders/test_doc_loader.py +++ b/openrag/services/workers/parsers/legacy_loaders/test_doc_loader.py @@ -12,7 +12,8 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest -from config.models import LoaderConfig, VLMConfig +from core.config.endpoints import VLMConfig +from core.config.indexation import LoaderConfig from core.models.document import ProcessedDocument, TextBlock from langchain_core.documents.base import Document as LCDocument diff --git a/openrag/services/workers/parsers/legacy_loaders/test_eml_recursion.py b/openrag/services/workers/parsers/legacy_loaders/test_eml_recursion.py index b1644096b..d69b9c641 100644 --- a/openrag/services/workers/parsers/legacy_loaders/test_eml_recursion.py +++ b/openrag/services/workers/parsers/legacy_loaders/test_eml_recursion.py @@ -87,7 +87,7 @@ async def test_eml_below_cap_does_not_skip(tmp_path): def test_recursion_cap_lives_in_loader_config(): """The cap must come from the loader config (operator-tunable), not be hardcoded in the loader. Assert the contract: a positive integer.""" - from config import load_config + from core.config import load_config cap = load_config().loader.eml_max_recursion_depth assert isinstance(cap, int) diff --git a/openrag/services/workers/parsers/marker_workers.py b/openrag/services/workers/parsers/marker_workers.py index 382cf6ec9..255d6fe4f 100644 --- a/openrag/services/workers/parsers/marker_workers.py +++ b/openrag/services/workers/parsers/marker_workers.py @@ -6,7 +6,7 @@ import pypdfium2 import ray import torch -from config import load_config +from core.config import load_config from core.indexing.image_preprocessor import pil_to_png_bytes from core.indexing.parsers.document_parser import BasePooledParser from core.models.document import ( @@ -33,7 +33,7 @@ class MarkerWorker: def __init__(self): import os - from config import load_config + from core.config import load_config from core.utils.logging import get_logger self.logger = get_logger() @@ -176,7 +176,7 @@ def __del__(self): @ray.remote(max_restarts=5) class MarkerPool: def __init__(self): - from config import load_config + from core.config import load_config from core.utils.logging import get_logger self.logger = get_logger() diff --git a/openrag/services/workers/parsers/whisper_workers.py b/openrag/services/workers/parsers/whisper_workers.py index ff2a959bf..990602d9a 100644 --- a/openrag/services/workers/parsers/whisper_workers.py +++ b/openrag/services/workers/parsers/whisper_workers.py @@ -3,7 +3,7 @@ import ray import torch -from config import load_config +from core.config import load_config from core.indexing.parsers.document_parser import BasePooledParser from core.models.document import ( Document, @@ -39,7 +39,7 @@ def whisper_actor_options(config) -> dict[str, float | int]: class WhisperActor: def __init__(self): import torch - from config import load_config + from core.config import load_config from core.utils.logging import get_logger self.logger = get_logger() @@ -102,7 +102,7 @@ class WhisperPool: """ def __init__(self): - from config import load_config + from core.config import load_config from core.utils.logging import get_logger self.logger = get_logger() diff --git a/openrag/services/workers/task_state.py b/openrag/services/workers/task_state.py index 9f2955a44..8f0584b70 100644 --- a/openrag/services/workers/task_state.py +++ b/openrag/services/workers/task_state.py @@ -7,7 +7,7 @@ import ray try: - from config import load_config as _load_config + from core.config import load_config as _load_config _cfg = _load_config() _POOL_SIZE: int = _cfg.ray.pool_size diff --git a/openrag/services/workers/test_indexer_pool.py b/openrag/services/workers/test_indexer_pool.py index 31129090d..3f9824b14 100644 --- a/openrag/services/workers/test_indexer_pool.py +++ b/openrag/services/workers/test_indexer_pool.py @@ -12,6 +12,10 @@ class _BrokenChunker: pass +class _NonCallableChunker: + chunk = None + + def test_build_chunker_returns_native_chunker(monkeypatch: pytest.MonkeyPatch) -> None: import core.chunking.factory as factory from services.workers.indexer_pool import _build_chunker @@ -30,3 +34,13 @@ def test_build_chunker_rejects_invalid_chunker(monkeypatch: pytest.MonkeyPatch) with pytest.raises(TypeError, match="chunk"): _build_chunker(object()) + + +def test_build_chunker_rejects_non_callable_chunk_attr(monkeypatch: pytest.MonkeyPatch) -> None: + import core.chunking.factory as factory + from services.workers.indexer_pool import _build_chunker + + monkeypatch.setattr(factory, "create_chunker", lambda _cfg: _NonCallableChunker()) + + with pytest.raises(TypeError, match="chunk"): + _build_chunker(object()) diff --git a/openrag/tests/test_relationships_integration.py b/openrag/tests/test_relationships_integration.py index 7b984e3c0..5c54df1ab 100644 --- a/openrag/tests/test_relationships_integration.py +++ b/openrag/tests/test_relationships_integration.py @@ -9,7 +9,7 @@ 2. Email thread scenario: Hierarchical email chain with parallel branches """ -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, MagicMock import pytest from langchain_core.documents.base import Document @@ -282,141 +282,3 @@ async def test_get_all_thread_emails_by_relationship(self, mock_vectordb_email, file_ids = {doc.metadata["file_id"] for doc in related} expected = {"email_original", "email_reply1", "email_reply2", "email_reply3"} assert file_ids == expected - - -class TestRelationshipAwareRetrieverIntegration: - """ - Integration tests for RelationshipAwareRetriever. - - Tests the retriever's ability to expand search results with - related and ancestor documents. - - Note: These tests require the configuration to be available. - They are marked to skip when the config is not found. - """ - - @pytest.fixture - def mock_documents(self): - """Create a set of related test documents.""" - return [ - Document( - page_content="Main document content", - metadata={ - "_id": "main_chunk", - "file_id": "main_file", - "partition": "test", - "relationship_id": "group_1", - "parent_id": None, - }, - ), - Document( - page_content="Related document 1", - metadata={ - "_id": "related_1", - "file_id": "related_file_1", - "partition": "test", - "relationship_id": "group_1", - "parent_id": "main_file", - }, - ), - Document( - page_content="Related document 2", - metadata={ - "_id": "related_2", - "file_id": "related_file_2", - "partition": "test", - "relationship_id": "group_1", - "parent_id": "main_file", - }, - ), - ] - - @pytest.mark.asyncio - async def test_retriever_without_expansion_returns_base_results(self): - """Test that retriever without expansion returns only base search results.""" - try: - from components.retriever import RelationshipAwareRetriever - except Exception: - pytest.skip("Requires config to be available") - - with patch("components.retriever.get_vectordb") as mock_get_db: - mock_db = MagicMock() - mock_db.async_search = MagicMock() - mock_db.async_search.remote = AsyncMock( - return_value=[Document(page_content="Test", metadata={"_id": "1", "partition": "test"})] - ) - mock_get_db.return_value = mock_db - - retriever = RelationshipAwareRetriever( - include_related=False, - include_ancestors=False, - ) - results = await retriever.retrieve(["test"], "query") - - assert len(results) == 1 - assert results[0].page_content == "Test" - - @pytest.mark.asyncio - async def test_retriever_with_include_related_expands_results(self, mock_documents): - """Test that retriever with include_related expands with related docs.""" - try: - from components.retriever import RelationshipAwareRetriever - except Exception: - pytest.skip("Requires config to be available") - - with patch("components.retriever.get_vectordb") as mock_get_db: - mock_db = MagicMock() - - # Base search returns main document - mock_db.async_search = MagicMock() - mock_db.async_search.remote = AsyncMock(return_value=[mock_documents[0]]) - - # Related chunks returns all related documents - mock_db.get_related_chunks = MagicMock() - mock_db.get_related_chunks.remote = AsyncMock(return_value=mock_documents) - - mock_get_db.return_value = mock_db - - retriever = RelationshipAwareRetriever( - include_related=True, - include_ancestors=False, - ) - results = await retriever.retrieve(["test"], "query") - - # Should have base result + related (deduplicated) - assert len(results) == 3 - chunk_ids = {r.metadata["_id"] for r in results} - assert chunk_ids == {"main_chunk", "related_1", "related_2"} - - @pytest.mark.asyncio - async def test_retriever_deduplicates_results(self, mock_documents): - """Test that retriever properly deduplicates expanded results.""" - try: - from components.retriever import RelationshipAwareRetriever - except Exception: - pytest.skip("Requires config to be available") - - with patch("components.retriever.get_vectordb") as mock_get_db: - mock_db = MagicMock() - - # Base search returns main document - mock_db.async_search = MagicMock() - mock_db.async_search.remote = AsyncMock(return_value=[mock_documents[0]]) - - # Related chunks returns same main document again (plus others) - mock_db.get_related_chunks = MagicMock() - mock_db.get_related_chunks.remote = AsyncMock(return_value=[mock_documents[0], mock_documents[1]]) - - mock_get_db.return_value = mock_db - - retriever = RelationshipAwareRetriever( - include_related=True, - include_ancestors=False, - ) - results = await retriever.retrieve(["test"], "query") - - # Should deduplicate, so only 2 unique results - assert len(results) == 2 - chunk_ids = [r.metadata["_id"] for r in results] - # Main chunk should appear only once - assert chunk_ids.count("main_chunk") == 1 diff --git a/openrag/utils/__init__.py b/openrag/utils/__init__.py deleted file mode 100644 index e69de29bb..000000000 diff --git a/openrag/utils/exceptions/__init__.py b/openrag/utils/exceptions/__init__.py deleted file mode 100644 index 7ed1e2a97..000000000 --- a/openrag/utils/exceptions/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -# Re-export from canonical location for backward compatibility. -# New code should import from `core.utils.exceptions` directly. -from core.utils.exceptions import * # noqa: F401,F403 diff --git a/openrag/utils/exceptions/base.py b/openrag/utils/exceptions/base.py deleted file mode 100644 index c0fed1294..000000000 --- a/openrag/utils/exceptions/base.py +++ /dev/null @@ -1,7 +0,0 @@ -# Re-export from canonical location for backward compatibility. -# New code should import from `core.utils.exceptions` directly. -from core.utils.exceptions import ( # noqa: F401 - EmbeddingError, - OpenRAGError, - VDBError, -) diff --git a/openrag/utils/exceptions/embeddings.py b/openrag/utils/exceptions/embeddings.py deleted file mode 100644 index 84be79cef..000000000 --- a/openrag/utils/exceptions/embeddings.py +++ /dev/null @@ -1,7 +0,0 @@ -# Re-export from canonical location for backward compatibility. -# New code should import from `core.utils.exceptions` directly. -from core.utils.exceptions import ( # noqa: F401 - EmbeddingAPIError, - EmbeddingResponseError, - UnexpectedEmbeddingError, -) diff --git a/openrag/utils/exceptions/vectordb.py b/openrag/utils/exceptions/vectordb.py deleted file mode 100644 index 54dd12df1..000000000 --- a/openrag/utils/exceptions/vectordb.py +++ /dev/null @@ -1,17 +0,0 @@ -# Re-export from canonical location for backward compatibility. -# New code should import from `core.utils.exceptions` directly. -from core.utils.exceptions import ( # noqa: F401 - UnexpectedVDBError, - VDBConnectionError, - VDBCreateOrLoadCollectionError, - VDBDeleteError, - VDBError, - VDBFileIDAlreadyExistsError, - VDBFileNotFoundError, - VDBInsertError, - VDBMembershipNotFound, - VDBPartitionNotFound, - VDBSchemaMigrationRequiredError, - VDBSearchError, - VDBUserNotFound, -) diff --git a/openrag/utils/external_resource_errors.py b/openrag/utils/external_resource_errors.py deleted file mode 100644 index 49f930a26..000000000 --- a/openrag/utils/external_resource_errors.py +++ /dev/null @@ -1,13 +0,0 @@ -"""Re-export from canonical location for backwards compatibility.""" - -from core.utils.external_errors import ( - EXTERNAL_ERROR_CODES, - EXTERNAL_ERROR_INDICATORS, - is_external_resource_error, -) - -__all__ = [ - "EXTERNAL_ERROR_CODES", - "EXTERNAL_ERROR_INDICATORS", - "is_external_resource_error", -] From 8c2beb40c730d005343fa5603a3bc785ddecad2e Mon Sep 17 00:00:00 2001 From: hedhoud <74668966+hedhoud@users.noreply.github.com> Date: Fri, 29 May 2026 16:46:09 +0200 Subject: [PATCH 11/11] test: remove stale legacy utils alias --- openrag/conftest.py | 10 ++++------ 1 file changed, 4 insertions(+), 6 deletions(-) diff --git a/openrag/conftest.py b/openrag/conftest.py index 7f0563d99..19f368e97 100644 --- a/openrag/conftest.py +++ b/openrag/conftest.py @@ -1,9 +1,8 @@ -"""Pytest import guards for legacy top-level imports. +"""Pytest import guards for top-level imports. -Some modules still import ``utils`` and the third-party ``openai`` package as -top-level names. During collection, pytest can prepend nested test directories -such as ``openrag/routers`` to ``sys.path``, where ``utils.py`` and -``openai.py`` would otherwise shadow those imports. +During collection, pytest can prepend nested test directories such as +``openrag/routers`` to ``sys.path``, where local modules could otherwise shadow +third-party imports such as ``openai``. """ from __future__ import annotations @@ -17,5 +16,4 @@ if root not in sys.path: sys.path.insert(0, root) -sys.modules.setdefault("utils", importlib.import_module("utils")) sys.modules.setdefault("openai", importlib.import_module("openai"))