diff --git a/.hydra_config/chunker/base.yaml b/.hydra_config/chunker/base.yaml deleted file mode 100644 index 08b13a60c..000000000 --- a/.hydra_config/chunker/base.yaml +++ /dev/null @@ -1,5 +0,0 @@ -contextual_retrieval: ${oc.decode:${oc.env:CONTEXTUAL_RETRIEVAL, true}} -contextualization_timeout: ${oc.decode:${oc.env:CONTEXTUALIZATION_TIMEOUT, 120}} -max_concurrent_contextualization: ${oc.decode:${oc.env:MAX_CONCURRENT_CONTEXTUALIZATION, 10}} -chunk_size: ${oc.decode:${oc.env:CHUNK_SIZE, 512}} -chunk_overlap_rate: ${oc.decode:${oc.env:CHUNK_OVERLAP_RATE, 0.2}} \ No newline at end of file diff --git a/.hydra_config/chunker/recursive_splitter.yaml b/.hydra_config/chunker/recursive_splitter.yaml deleted file mode 100644 index dc16946c4..000000000 --- a/.hydra_config/chunker/recursive_splitter.yaml +++ /dev/null @@ -1,8 +0,0 @@ -defaults: - - base -name: recursive_splitter - - -# https://chat.deepseek.com/a/chat/s/28913c5d-1f62-40b0-9247-4655994fe16b -# 3000 c => 750 tokens => -# 1 p => 2000 c => 500 tokens \ No newline at end of file diff --git a/.hydra_config/config.yaml b/.hydra_config/config.yaml deleted file mode 100644 index e7e30cb95..000000000 --- a/.hydra_config/config.yaml +++ /dev/null @@ -1,196 +0,0 @@ -defaults: - - _self_ # TODO: Silences the hydra version migration warning (PLEASE REVIEW FOR BREAKING CHANGES) - - chunker: ${oc.env:CHUNKER, recursive_splitter} # recursive_splitter - - retriever: ${oc.env:RETRIEVER_TYPE, single} # single # multiQuery # hyde - - rag: ChatBotRag - - websearch: ${oc.env:WEBSEARCH_PROVIDER, staan} - -llm_params: &llm_params - temperature: 0.1 - timeout: 60 - max_retries: 2 - logprobs: true - -llm: - <<: *llm_params - base_url: ${oc.env:BASE_URL} - model: ${oc.env:MODEL} - api_key: ${oc.env:API_KEY} - -vlm: - <<: *llm_params - base_url: ${oc.env:VLM_BASE_URL} - model: ${oc.env:VLM_MODEL} - api_key: ${oc.env:VLM_API_KEY} - -semaphore: - llm_semaphore: ${oc.decode:${oc.env:LLM_SEMAPHORE, 10}} - vlm_semaphore: ${oc.decode:${oc.env:VLM_SEMAPHORE, 10}} - -embedder: - provider: openai - model_name: ${oc.env:EMBEDDER_MODEL_NAME, jinaai/jina-embeddings-v3} - base_url: ${oc.env:EMBEDDER_BASE_URL, http://vllm:8000/v1} - api_key: ${oc.env:EMBEDDER_API_KEY, EMPTY} - max_model_len: ${oc.decode:${oc.env:MAX_MODEL_LEN, 8192}} - -vectordb: - host: ${oc.env:VDB_HOST, milvus} - port: ${oc.env:VDB_iPORT, 19530} - connector_name: ${oc.env:VDB_CONNECTOR_NAME, milvus} - collection_name: ${oc.env:VDB_COLLECTION_NAME, vdb_test} - hybrid_search: ${oc.env:VDB_HYBRID_SEARCH, true} - enable: true - -rdb: - host: ${oc.env:POSTGRES_HOST, rdb} - port: ${oc.env:POSTGRES_PORT, 5432} - user: ${oc.env:POSTGRES_USER, root} - password: ${oc.env:POSTGRES_PASSWORD, root_password} - default_file_quota: ${oc.decode:${oc.env:DEFAULT_FILE_QUOTA, -1}} - -reranker: - enable: ${oc.decode:${oc.env:RERANKER_ENABLED, true}} - model_name: ${oc.env:RERANKER_MODEL, Alibaba-NLP/gte-multilingual-reranker-base} - top_k: ${oc.decode:${oc.env:RERANKER_TOP_K, 10}} # Number of documents to return after reranking. Upgrade for better results if your llm has a wider context window. - base_url: ${oc.env:RERANKER_BASE_URL, http://reranker:${oc.env:RERANKER_PORT, 7997}} - -map_reduce: - # Number of documents to process in the initial mapping phase - initial_batch_size: ${oc.decode:${oc.env:MAP_REDUCE_INITIAL_BATCH_SIZE, 10}} - - # Number of additional documents to probe when all previous chunks are relevant - expansion_batch_size: ${oc.decode:${oc.env:MAP_REDUCE_EXPANSION_BATCH_SIZE, 5}} - - # Maximum total number of documents (chunks) to process across all iterations - max_total_documents: ${oc.decode:${oc.env:MAP_REDUCE_MAX_TOTAL_DOCUMENTS, 20}} - - # Enable debug logging for map & reduce - debug: ${oc.decode:${oc.env:MAP_REDUCE_DEBUG, false}} - - -verbose: - level: ${oc.env:LOG_LEVEL, DEBUG} - -server: - preferred_url_scheme: ${oc.env:PREFERRED_URL_SCHEME, null} - -llm_context: - max_llm_context_size: ${oc.decode:${oc.env:MAX_LLM_CONTEXT_SIZE, 8192}} - max_output_tokens: ${oc.decode:${oc.env:MAX_OUTPUT_TOKENS, 1024}} - -paths: - prompts_dir: ${oc.env:PROMPTS_DIR, ../prompts/example1} - data_dir: ${oc.env:DATA_DIR, ../data} - db_dir: ${oc.env:DB_DIR, /app/db} - log_dir: ${oc.env:LOG_DIR, /app/logs} - -prompts: - sys_prompt: sys_prompt_tmpl.txt - query_contextualizer: query_contextualizer_tmpl.txt - chunk_contextualizer: chunk_contextualizer_tmpl.txt - image_describer: image_captioning_tmpl.txt - spoken_style_answer: spoken_style_answer_tmpl.txt - - # query templates for different retriever types - hyde: hyde.txt - multi_query: multi_query_pmpt_tmpl.txt - -loader: - image_captioning: ${oc.decode:${oc.env:IMAGE_CAPTIONING, true}} - image_captioning_url: ${oc.decode:${oc.env:IMAGE_CAPTIONING_URL, true}} - save_markdown: ${oc.decode:${oc.env:SAVE_MARKDOWN, false}} - mimetypes: - text/plain: .txt - text/markdown: .md - application/pdf: .pdf - message/rfc822: .eml - application/vnd.openxmlformats-officedocument.wordprocessingml.document: .docx - application/vnd.openxmlformats-officedocument.presentationml.presentation: .pptx - application/msword: .doc - image/png: .png - image/jpeg: .jpeg - audio/wav: .wav - audio/mpeg: .mp3 - audio/flac: .flac - audio/ogg: .ogg - audio/aac: .aac - video/x-flv: .flv - audio/x-ms-wma: .wma - video/mp4: .mp4 - local_whisper: - model: ${oc.env:WHISPER_MODEL, base} # large-v3-turbo, base, small, medium, large - whisper_n_workers: ${oc.decode:${oc.env:WHISPER_N_WORKERS, 3}} # Number of parallel whisper workers / actors to launch - whisper_num_gpus: ${oc.decode:${oc.env:WHISPER_NUM_GPUS, 0.01}} - whisper_concurency_per_worker : ${oc.decode:${oc.env:WHISPER_CONCURRENCY_PER_WORKER, 2}} # Number of concurrent transcriptions per whisper worker / actor. Increase if you have many small files and enough CPU/GPU resources. - file_loaders: - txt: TextLoader - pdf: ${oc.env:PDFLoader, MarkerLoader} # DoclingLoader # MarkerLoader # PyMuPDFLoader # Custompymupdf4llm - eml: EmlLoader - docx: DocxLoader - pptx: PPTXLoader - doc: DocLoader - png: ImageLoader - jpeg: ImageLoader - jpg: ImageLoader - svg: ImageLoader - # Audio formats - wav: ${oc.env:AUDIOLOADER, LocalWhisperLoader} # LocalWhisperLoader # OpenAIAUDIOLOADER - mp3: ${oc.env:AUDIOLOADER, LocalWhisperLoader} - flac: ${oc.env:AUDIOLOADER, LocalWhisperLoader} - ogg: ${oc.env:AUDIOLOADER, LocalWhisperLoader} - aac: ${oc.env:AUDIOLOADER, LocalWhisperLoader} - flv: ${oc.env:AUDIOLOADER, LocalWhisperLoader} - wma: ${oc.env:AUDIOLOADER, LocalWhisperLoader} - # Video formats - mp4: ${oc.env:AUDIOLOADER, LocalWhisperLoader} - md: MarkdownLoader - marker_max_tasks_per_child: ${oc.decode:${oc.env:MARKER_MAX_TASKS_PER_CHILD, 10}} - marker_pool_size: ${oc.decode:${oc.env:MARKER_POOL_SIZE, 1}} # Value au increment if you have a cluster of machines - marker_max_processes: ${oc.decode:${oc.env:MARKER_MAX_PROCESSES, 2}} # Maximum number of concurrent PDF parsing processes - marker_min_processes: ${oc.decode:${oc.env:MARKER_MIN_PROCESSES, 1}} - marker_num_gpus: ${oc.decode:${oc.env:MARKER_NUM_GPUS, 0.01}} # See https://docs.ray.io/en/latest/cluster/kubernetes/user-guides/gpu.html - marker_timeout: ${oc.decode:${oc.env:MARKER_TIMEOUT, 3600}} - marker_pdftext_workers: ${oc.decode:${oc.env:MARKER_PDFTEXT_WORKERS, 2}} - - transcriber: - base_url: ${oc.env:TRANSCRIBER_BASE_URL, http://transcriber:8000/v1} - api_key: ${oc.env:TRANSCRIBER_API_KEY, EMPTY} - model_name: ${oc.env:TRANSCRIBER_MODEL, openai/whisper-large-v3-turbo} - timeout: ${oc.decode:${oc.env:TRANSCRIBER_TIMEOUT, 3600}} - max_concurrent_chunks: ${oc.decode:${oc.env:TRANSCRIBER_MAX_CONCURRENT_CHUNKS, 20}} - use_whisper_lang_detector: ${oc.decode:${oc.env:USE_WHISPER_LANG_DETECTOR, true}} - openai: - base_url: ${oc.env:OPENAI_LOADER_BASE_URL, http://openai:8000/v1} - api_key: ${oc.env:OPENAI_LOADER_API_KEY, EMPTY} - model: ${oc.env:OPENAI_LOADER_MODEL, dotsocr-model} - temperature: ${oc.decode:${oc.env:OPENAI_LOADER_TEMPERATURE, 0.2}} - timeout: ${oc.decode:${oc.env:OPENAI_LOADER_TIMEOUT, 180}} - max_retries: ${oc.decode:${oc.env:OPENAI_LOADER_MAX_RETRIES, 2}} - top_p: ${oc.decode:${oc.env:OPENAI_LOADER_TOP_P, 0.9}} - concurrency_limit: ${oc.decode:${oc.env:OPENAI_LOADER_CONCURRENCY_LIMIT, 20}} - -ray: - num_gpus: ${oc.decode:${oc.env:RAY_NUM_GPUS, 0.01}} - pool_size: ${oc.decode:${oc.env:RAY_POOL_SIZE, 1}} # Number of serializer actor instances - max_tasks_per_worker: ${oc.decode:${oc.env:RAY_MAX_TASKS_PER_WORKER, 8}} # Number of tasks per serializer instance - indexer: - max_task_retries: ${oc.decode:${oc.env:RAY_MAX_TASK_RETRIES, 2}} - serialize_timeout: ${oc.decode:${oc.env:INDEXER_SERIALIZE_TIMEOUT, 3600}} - vectordb_timeout: ${oc.decode:${oc.env:VECTORDB_TIMEOUT, 30}} - concurrency_groups: - default: ${oc.decode:${oc.env:INDEXER_DEFAULT_CONCURRENCY, 1000}} - update: ${oc.decode:${oc.env:INDEXER_UPDATE_CONCURRENCY, 100}} - search: ${oc.decode:${oc.env:INDEXER_SEARCH_CONCURRENCY, 100}} - delete: ${oc.decode:${oc.env:INDEXER_DELETE_CONCURRENCY, 100}} - serialize: ${oc.decode:${oc.env:INDEXER_SERIALIZE_CONCURRENCY, 50}} - chunk: ${oc.decode:${oc.env:INDEXER_CHUNK_CONCURRENCY, 1000}} - insert: ${oc.decode:${oc.env:INDEXER_INSERT_CONCURRENCY, 100}} - semaphore: - concurrency: ${oc.decode:${oc.env:RAY_SEMAPHORE_CONCURRENCY, 100000}} - serve: - enable: ${oc.decode:${oc.env:ENABLE_RAY_SERVE, false}} - num_replicas: ${oc.decode:${oc.env:RAY_SERVE_NUM_REPLICAS, 1}} - host: ${oc.env:RAY_SERVE_HOST, 0.0.0.0} - port: ${oc.env:RAY_SERVE_PORT, 8080} - chainlit_port: ${oc.env:CHAINLIT_PORT, 8090} \ No newline at end of file diff --git a/.hydra_config/rag/ChatBotRag.yaml b/.hydra_config/rag/ChatBotRag.yaml deleted file mode 100644 index 40f8165de..000000000 --- a/.hydra_config/rag/ChatBotRag.yaml +++ /dev/null @@ -1,3 +0,0 @@ -defaults: - - base -mode: ChatBotRag \ No newline at end of file diff --git a/.hydra_config/rag/SimpleRag.yaml b/.hydra_config/rag/SimpleRag.yaml deleted file mode 100644 index 59344a711..000000000 --- a/.hydra_config/rag/SimpleRag.yaml +++ /dev/null @@ -1,3 +0,0 @@ -defaults: - - base -mode: SimpleRag \ No newline at end of file diff --git a/.hydra_config/rag/base.yaml b/.hydra_config/rag/base.yaml deleted file mode 100644 index 3a09f558c..000000000 --- a/.hydra_config/rag/base.yaml +++ /dev/null @@ -1,4 +0,0 @@ -# Config for chatbot RAG -mode: '' -chat_history_depth: 4 -max_contextualized_query_len: 512 \ No newline at end of file diff --git a/.hydra_config/retriever/base.yaml b/.hydra_config/retriever/base.yaml deleted file mode 100644 index 13ec22748..000000000 --- a/.hydra_config/retriever/base.yaml +++ /dev/null @@ -1,8 +0,0 @@ -type: '' -top_k: ${oc.decode:${oc.env:RETRIEVER_TOP_K, 50}} # Number of documents to return before reranking -similarity_threshold: ${oc.decode:${oc.env:SIMILARITY_THRESHOLD, 0.6}} # Minimum similarity score for document retrieval -with_surrounding_chunks: ${oc.decode:${oc.env:WITH_SURROUNDING_CHUNKS, false}} # Whether to include surrounding chunks for each retrieved chunk -include_related: ${oc.decode:${oc.env:INCLUDE_RELATED, true}} -include_ancestors: ${oc.decode:${oc.env:INCLUDE_ANCESTORS, true}} -related_limit: ${oc.decode:${oc.env:RELATED_LIMIT, 10}} -max_ancestor_depth: ${oc.decode:${oc.env:MAX_DEPTH, 10}} # Maximum depth for ancestor retrieval (null = unlimited) \ No newline at end of file diff --git a/.hydra_config/retriever/hyde.yaml b/.hydra_config/retriever/hyde.yaml deleted file mode 100644 index 6a246638d..000000000 --- a/.hydra_config/retriever/hyde.yaml +++ /dev/null @@ -1,6 +0,0 @@ -defaults: - - base - -type: hyde -# Extra params -combine: False \ No newline at end of file diff --git a/.hydra_config/retriever/multiQuery.yaml b/.hydra_config/retriever/multiQuery.yaml deleted file mode 100644 index bc8ab73cf..000000000 --- a/.hydra_config/retriever/multiQuery.yaml +++ /dev/null @@ -1,5 +0,0 @@ -defaults: - - base - -type: multiQuery -k_queries: 3 \ No newline at end of file diff --git a/.hydra_config/retriever/single.yaml b/.hydra_config/retriever/single.yaml deleted file mode 100644 index 6be61d766..000000000 --- a/.hydra_config/retriever/single.yaml +++ /dev/null @@ -1,4 +0,0 @@ -defaults: - - base - -type: single \ No newline at end of file diff --git a/.hydra_config/websearch/base.yaml b/.hydra_config/websearch/base.yaml deleted file mode 100644 index 83381dd81..000000000 --- a/.hydra_config/websearch/base.yaml +++ /dev/null @@ -1,11 +0,0 @@ -provider: '' -api_token: ${oc.env:WEBSEARCH_API_TOKEN, ""} -base_url: '' -top_k: ${oc.decode:${oc.env:WEBSEARCH_TOP_K, 5}} -lang: ${oc.env:WEBSEARCH_LANG, fr-FR} -max_tokens: ${oc.decode:${oc.env:WEBSEARCH_MAX_TOKENS, 2000}} -fetch_content: ${oc.decode:${oc.env:WEBSEARCH_FETCH_CONTENT, true}} -fetch_max_results: ${oc.decode:${oc.env:WEBSEARCH_FETCH_MAX_RESULTS, 3}} -fetch_timeout: ${oc.decode:${oc.env:WEBSEARCH_FETCH_TIMEOUT, 1.0}} -fetch_max_tokens: ${oc.decode:${oc.env:WEBSEARCH_FETCH_MAX_TOKENS, 500}} -fetch_verify_ssl: ${oc.decode:${oc.env:WEBSEARCH_FETCH_VERIFY_SSL, false}} diff --git a/.hydra_config/websearch/staan.yaml b/.hydra_config/websearch/staan.yaml deleted file mode 100644 index de14bc331..000000000 --- a/.hydra_config/websearch/staan.yaml +++ /dev/null @@ -1,5 +0,0 @@ -defaults: - - base - -provider: staan -base_url: ${oc.env:WEBSEARCH_BASE_URL, "https://api.staan.ai/search/web"} diff --git a/Dockerfile b/Dockerfile index 75262c3c6..01e35d3d4 100644 --- a/Dockerfile +++ b/Dockerfile @@ -41,9 +41,9 @@ WORKDIR /app/openrag # Copy source code COPY openrag/ . -# Copy assests & config +# Copy assets and config COPY prompts/ /app/prompts/ -COPY .hydra_config/ /app/.hydra_config/ +COPY conf/ /app/conf/ ENV PYTHONPATH=/app/openrag/ ENV APP_iPORT=${APP_iPORT:-8080} ENTRYPOINT ../entrypoint.sh diff --git a/Dockerfile.ray b/Dockerfile.ray index 8f91d86c7..f56ea6370 100644 --- a/Dockerfile.ray +++ b/Dockerfile.ray @@ -46,11 +46,10 @@ WORKDIR /app/openrag # Copy source code COPY openrag/ . -# Copy assests & config +# Copy assets and config COPY prompts/ /app/prompts/ -COPY .hydra_config/ /app/.hydra_config/ +COPY conf/ /app/conf/ -RUN apt install -y RUN ln -s /app/.venv/bin/ray /usr/local/bin/ray ENV PYTHONPATH=/app/openrag/ \ No newline at end of file diff --git a/cluster.yaml b/cluster.yaml index 92b06fec6..066c9c22c 100644 --- a/cluster.yaml +++ b/cluster.yaml @@ -12,7 +12,6 @@ docker: - --gpus all - -v /ray_mount/model_weights:/app/model_weights - -v /ray_mount/data:/app/data - - -v /ray_mount/.hydra_config:/app/.hydra_config - -v /ray_mount/logs:/app/logs - --env-file /ray_mount/.env auth: diff --git a/conf/config.yaml b/conf/config.yaml new file mode 100644 index 000000000..cfe2da83c --- /dev/null +++ b/conf/config.yaml @@ -0,0 +1,297 @@ +# ============================================================================= +# OpenRAG Configuration +# ============================================================================= +# This file is the single source of truth for all default settings. +# Values can be overridden via environment variables (see comments). +# +# Secrets (API keys, passwords) should NEVER be set here β€” use env vars +# or a .env file instead. +# ============================================================================= + +# --- Shared LLM parameters (YAML anchor, not a config section) --- +_llm_params: &llm_params + temperature: 0.1 + timeout: 60 + max_retries: 2 + logprobs: true + +# --- LLM --- +# Env: BASE_URL, MODEL, API_KEY +llm: + <<: *llm_params + base_url: "" + model: "" + api_key: "" + +# --- VLM (Vision Language Model) --- +# Env: VLM_BASE_URL, VLM_MODEL, VLM_API_KEY +vlm: + <<: *llm_params + base_url: "" + model: "" + api_key: "" + +# --- Semaphore (concurrency limits) --- +# Env: LLM_SEMAPHORE, VLM_SEMAPHORE +semaphore: + llm_semaphore: 10 + vlm_semaphore: 10 + +# --- Embedder --- +# Env: EMBEDDER_MODEL_NAME, EMBEDDER_BASE_URL, EMBEDDER_API_KEY, MAX_MODEL_LEN +embedder: + provider: openai + model_name: jinaai/jina-embeddings-v3 + base_url: http://vllm:8000/v1 + api_key: EMPTY + max_model_len: 8192 + +# --- Vector Database (Milvus) --- +# Env: VDB_HOST, VDB_PORT, VDB_CONNECTOR_NAME, VDB_COLLECTION_NAME, VDB_HYBRID_SEARCH +vectordb: + host: milvus + port: 19530 + connector_name: milvus + collection_name: vdb_test + hybrid_search: true + enable: true + +# --- Relational Database (PostgreSQL) --- +# Env: POSTGRES_HOST, POSTGRES_PORT, POSTGRES_USER, POSTGRES_PASSWORD, DEFAULT_FILE_QUOTA +rdb: + host: rdb + port: 5432 + user: root + password: "root_password" + default_file_quota: -1 + +# --- Reranker --- +# Env: RERANKER_ENABLED, RERANKER_MODEL, RERANKER_TOP_K, RERANKER_BASE_URL, RERANKER_PORT +reranker: + enable: true + model_name: Alibaba-NLP/gte-multilingual-reranker-base + top_k: 10 + base_url: "" # Default built from RERANKER_PORT if empty + +# --- Map-Reduce --- +# Env: MAP_REDUCE_INITIAL_BATCH_SIZE, MAP_REDUCE_EXPANSION_BATCH_SIZE, +# MAP_REDUCE_MAX_TOTAL_DOCUMENTS, MAP_REDUCE_DEBUG +map_reduce: + initial_batch_size: 10 + expansion_batch_size: 5 + max_total_documents: 20 + debug: false + +# --- Logging --- +# Env: LOG_LEVEL +verbose: + level: DEBUG + +# --- Server --- +# Env: PREFERRED_URL_SCHEME +server: + preferred_url_scheme: null + +# --- LLM Context --- +# Env: MAX_LLM_CONTEXT_SIZE, MAX_OUTPUT_TOKENS +llm_context: + max_llm_context_size: 8192 + max_output_tokens: 1024 + +# --- Paths --- +# Env: PROMPTS_DIR, DATA_DIR, DB_DIR, LOG_DIR +paths: + prompts_dir: ../prompts/example1 + data_dir: ../data + db_dir: /app/db + log_dir: /app/logs + +# --- Prompt template filenames --- +prompts: + sys_prompt: sys_prompt_tmpl.txt + query_contextualizer: query_contextualizer_tmpl.txt + chunk_contextualizer: chunk_contextualizer_tmpl.txt + image_describer: image_captioning_tmpl.txt + spoken_style_answer: spoken_style_answer_tmpl.txt + hyde: hyde.txt + multi_query: multi_query_pmpt_tmpl.txt + +# --- Document loader --- +loader: + # Env: IMAGE_CAPTIONING, IMAGE_CAPTIONING_URL, SAVE_MARKDOWN + image_captioning: true + image_captioning_url: true + save_markdown: false + + mimetypes: + text/plain: .txt + text/markdown: .md + application/pdf: .pdf + message/rfc822: .eml + application/vnd.openxmlformats-officedocument.wordprocessingml.document: .docx + application/vnd.openxmlformats-officedocument.presentationml.presentation: .pptx + application/msword: .doc + image/png: .png + image/jpeg: .jpeg + audio/wav: .wav + audio/mpeg: .mp3 + audio/flac: .flac + audio/ogg: .ogg + audio/aac: .aac + video/x-flv: .flv + audio/x-ms-wma: .wma + video/mp4: .mp4 + + # Env: WHISPER_MODEL, WHISPER_N_WORKERS, WHISPER_NUM_GPUS, WHISPER_CONCURRENCY_PER_WORKER + local_whisper: + model: base + whisper_n_workers: 3 + whisper_num_gpus: 0.01 + whisper_concurrency_per_worker: 2 + + # Env: PDFLoader, AUDIOLOADER (per-extension overrides) + file_loaders: + txt: TextLoader + pdf: MarkerLoader + eml: EmlLoader + docx: DocxLoader + pptx: PPTXLoader + doc: DocLoader + png: ImageLoader + jpeg: ImageLoader + jpg: ImageLoader + svg: ImageLoader + wav: LocalWhisperLoader + mp3: LocalWhisperLoader + flac: LocalWhisperLoader + ogg: LocalWhisperLoader + aac: LocalWhisperLoader + flv: LocalWhisperLoader + wma: LocalWhisperLoader + mp4: LocalWhisperLoader + md: MarkdownLoader + + # Env: MARKER_MAX_TASKS_PER_CHILD, MARKER_POOL_SIZE, MARKER_MAX_PROCESSES, + # MARKER_MIN_PROCESSES, MARKER_NUM_GPUS, MARKER_TIMEOUT, MARKER_PDFTEXT_WORKERS + marker_max_tasks_per_child: 10 + marker_pool_size: 1 + marker_max_processes: 2 + marker_min_processes: 1 + marker_num_gpus: 0.01 + marker_timeout: 3600 + marker_pdftext_workers: 2 + + # Env: DOCLING_NUM_GPUS, DOCLING_POOL_SIZE, DOCLING_MAX_TASKS_PER_WORKER + docling_num_gpus: 0.01 + docling_pool_size: 1 + docling_max_tasks_per_worker: 2 + + # Env: TRANSCRIBER_BASE_URL, TRANSCRIBER_API_KEY, TRANSCRIBER_MODEL, + # TRANSCRIBER_TIMEOUT, TRANSCRIBER_MAX_CONCURRENT_CHUNKS, USE_WHISPER_LANG_DETECTOR + transcriber: + base_url: http://transcriber:8000/v1 + api_key: EMPTY + model_name: openai/whisper-large-v3-turbo + timeout: 3600 + max_concurrent_chunks: 20 + use_whisper_lang_detector: true + + # Env: OPENAI_LOADER_BASE_URL, OPENAI_LOADER_API_KEY, OPENAI_LOADER_MODEL, + # OPENAI_LOADER_TEMPERATURE, OPENAI_LOADER_TIMEOUT, OPENAI_LOADER_MAX_RETRIES, + # OPENAI_LOADER_TOP_P, OPENAI_LOADER_CONCURRENCY_LIMIT + openai: + base_url: http://openai:8000/v1 + api_key: EMPTY + model: dotsocr-model + temperature: 0.2 + timeout: 180 + max_retries: 2 + top_p: 0.9 + concurrency_limit: 20 + +# --- Ray --- +ray: + # Env: RAY_NUM_GPUS, RAY_POOL_SIZE, RAY_MAX_TASKS_PER_WORKER + num_gpus: 0.01 + pool_size: 1 + max_tasks_per_worker: 8 + + indexer: + # Env: RAY_MAX_TASK_RETRIES, INDEXER_SERIALIZE_TIMEOUT, VECTORDB_TIMEOUT + max_task_retries: 2 + serialize_timeout: 3600 + vectordb_timeout: 30 + concurrency_groups: + # Env: INDEXER_DEFAULT_CONCURRENCY, INDEXER_UPDATE_CONCURRENCY, etc. + default: 1000 + update: 100 + search: 100 + delete: 100 + serialize: 50 + chunk: 1000 + insert: 100 + + semaphore: + # Env: RAY_SEMAPHORE_CONCURRENCY + concurrency: 100000 + + serve: + # Env: ENABLE_RAY_SERVE, RAY_SERVE_NUM_REPLICAS, RAY_SERVE_HOST, + # RAY_SERVE_PORT, CHAINLIT_PORT + enable: false + num_replicas: 1 + host: 0.0.0.0 + port: 8080 + chainlit_port: 8090 + +# --- Chunker --- +# Env: CHUNKER, CONTEXTUAL_RETRIEVAL, CONTEXTUALIZATION_TIMEOUT, +# MAX_CONCURRENT_CONTEXTUALIZATION, CHUNK_SIZE, CHUNK_OVERLAP_RATE +chunker: + name: recursive_splitter + contextual_retrieval: true + contextualization_timeout: 120 + max_concurrent_contextualization: 10 + chunk_size: 512 + chunk_overlap_rate: 0.2 + +# --- Retriever --- +# Env: RETRIEVER_TYPE, RETRIEVER_TOP_K, SIMILARITY_THRESHOLD, +# WITH_SURROUNDING_CHUNKS, INCLUDE_RELATED, INCLUDE_ANCESTORS, +# RELATED_LIMIT, MAX_DEPTH +retriever: + type: single + top_k: 50 + similarity_threshold: 0.6 + with_surrounding_chunks: false + include_related: true + include_ancestors: true + related_limit: 10 + max_ancestor_depth: 10 + k_queries: 3 + combine: false + +# --- RAG --- +# Env: RAG_MODE +rag: + mode: ChatBotRag + chat_history_depth: 4 + max_contextualized_query_len: 512 + +# --- Web Search --- +# Env: WEBSEARCH_PROVIDER, WEBSEARCH_API_TOKEN, WEBSEARCH_BASE_URL, +# WEBSEARCH_TOP_K, WEBSEARCH_LANG, WEBSEARCH_MAX_TOKENS, +# WEBSEARCH_FETCH_CONTENT, WEBSEARCH_FETCH_MAX_RESULTS, +# WEBSEARCH_FETCH_TIMEOUT, WEBSEARCH_FETCH_MAX_TOKENS, WEBSEARCH_FETCH_VERIFY_SSL +websearch: + provider: staan + api_token: "" + base_url: https://api.staan.ai/search/web + top_k: 5 + lang: fr-FR + max_tokens: 2000 + fetch_content: true + fetch_max_results: 3 + fetch_timeout: 1.0 + fetch_max_tokens: 500 + fetch_verify_ssl: false diff --git a/docker-compose.yaml b/docker-compose.yaml index 58dd0f93a..42a09354b 100644 --- a/docker-compose.yaml +++ b/docker-compose.yaml @@ -10,7 +10,6 @@ x-openrag: &openrag_template context: . dockerfile: Dockerfile volumes: - - ${CONFIG_VOLUME:-./.hydra_config}:/app/.hydra_config # For dev mode - ${DATA_VOLUME:-./data}:/app/data - ${MODEL_WEIGHTS_VOLUME:-~/.cache/huggingface}:/app/model_weights # Model weights for RAG - ./openrag:/app/openrag # For dev mode diff --git a/docs/assets/compose_linux_gpu.yaml b/docs/assets/compose_linux_gpu.yaml index ee75c64be..f13d68cf7 100644 --- a/docs/assets/compose_linux_gpu.yaml +++ b/docs/assets/compose_linux_gpu.yaml @@ -9,7 +9,6 @@ x-openrag: &openrag_template context: . dockerfile: Dockerfile volumes: - - ${CONFIG_VOLUME:-./.hydra_config}:/app/.hydra_config - ${DATA_VOLUME:-./data}:/app/data - ${MODEL_WEIGHTS_VOLUME:-~/.cache/huggingface}:/app/model_weights # Model weights for RAG - ./openrag:/app/openrag # For dev mode diff --git a/docs/content/docs/documentation/API.mdx b/docs/content/docs/documentation/API.mdx index 191f02662..c915b1ad9 100644 --- a/docs/content/docs/documentation/API.mdx +++ b/docs/content/docs/documentation/API.mdx @@ -55,6 +55,18 @@ Get openRAG version GET /version ``` +### βš™οΈ Configuration + +Get the current application configuration. Sensitive fields (`api_key`, `password`, `token`) are redacted. + +```http +GET /config +``` + +**Permissions:** Requires admin role + +**Response:** JSON object with all configuration sections (LLM, embedder, vector DB, chunker, retriever, etc.) + --- ### πŸ“¦ Document Indexing diff --git a/openrag/api.py b/openrag/api.py index 768fcf691..13c85d70a 100644 --- a/openrag/api.py +++ b/openrag/api.py @@ -10,7 +10,7 @@ import uvicorn from config import load_config from dotenv import dotenv_values -from fastapi import FastAPI, Request +from fastapi import Depends, FastAPI, Request from fastapi.middleware.cors import CORSMiddleware from fastapi.openapi.utils import get_openapi from fastapi.responses import JSONResponse @@ -32,6 +32,7 @@ from routers.search import router as search_router from routers.tools import router as tools_router from routers.users import router as users_router +from routers.utils import require_admin from routers.workspaces import router as workspaces_router from starlette.middleware.base import BaseHTTPMiddleware from utils.dependencies import get_vectordb @@ -237,6 +238,11 @@ def get_version(): return {"version": app.version} +@app.get("/config", summary="Get current configuration", tags=["Configuration"], dependencies=[Depends(require_admin)]) +def get_config(): + return config + + # Mount the indexer router app.include_router(indexer_router, prefix="/indexer", tags=[Tags.INDEXER]) # Mount the extract router diff --git a/openrag/components/indexer/chunker/chunker.py b/openrag/components/indexer/chunker/chunker.py index 3a8e5f1d1..de9f8934c 100644 --- a/openrag/components/indexer/chunker/chunker.py +++ b/openrag/components/indexer/chunker/chunker.py @@ -7,7 +7,6 @@ from langchain_core.documents.base import Document from langchain_core.messages import HumanMessage, SystemMessage from langchain_openai import ChatOpenAI -from omegaconf import OmegaConf from tqdm.asyncio import tqdm from utils.logger import get_logger @@ -20,9 +19,9 @@ config = load_config() # Timeout for individual chunk contextualization LLM calls (in seconds) -CONTEXTUALIZATION_TIMEOUT = config.chunker.get("contextualization_timeout", 120) +CONTEXTUALIZATION_TIMEOUT = config.chunker.contextualization_timeout # Maximum concurrent contextualization tasks to prevent system overload -MAX_CONCURRENT_CONTEXTUALIZATION = config.chunker.get("max_concurrent_contextualization", 10) +MAX_CONCURRENT_CONTEXTUALIZATION = config.chunker.max_concurrent_contextualization BASE_CHUNK_FORMAT = "* filename: {filename}\n\n[CHUNK_START]\n\n{content}\n\n[CHUNK_END]" CHUNK_FORMAT = "[CONTEXT]\n\n{chunk_context}\n\n" + BASE_CHUNK_FORMAT @@ -343,11 +342,11 @@ class ChunkerFactory: @staticmethod def create_chunker( - config: OmegaConf, + config, embedder: BaseEmbedding | None = None, ) -> BaseChunker: # Extract parameters - chunker_params = OmegaConf.to_container(config.chunker, resolve=True) + chunker_params = config.chunker.model_dump() name = chunker_params.pop("name") # Initialize and return the chunker @@ -358,5 +357,5 @@ def create_chunker( f"Chunker '{name}' is not recognized. Available chunkers: {list(ChunkerFactory.CHUNKERS.keys())}" ) - chunker_params["llm_config"] = config.vlm + chunker_params["llm_config"] = config.vlm.model_dump() return chunker_cls(**chunker_params) diff --git a/openrag/components/indexer/embeddings/__init__.py b/openrag/components/indexer/embeddings/__init__.py index 110056779..69850e74f 100644 --- a/openrag/components/indexer/embeddings/__init__.py +++ b/openrag/components/indexer/embeddings/__init__.py @@ -8,8 +8,8 @@ class EmbeddingFactory: @staticmethod - def get_embedder(embeddings_config: dict) -> BaseEmbedding: - provider = embeddings_config.get("provider") + def get_embedder(embeddings_config) -> BaseEmbedding: + provider = embeddings_config.provider embedder_class = EMBEDDER_MAPPING.get(provider, None) if not embedder_class: diff --git a/openrag/components/indexer/embeddings/openai.py b/openrag/components/indexer/embeddings/openai.py index 9dc031e74..23aafd904 100644 --- a/openrag/components/indexer/embeddings/openai.py +++ b/openrag/components/indexer/embeddings/openai.py @@ -10,11 +10,11 @@ class OpenAIEmbedding(BaseEmbedding): - def __init__(self, embeddings_config: dict): - self.embedding_model = embeddings_config.get("model_name") - self.base_url = embeddings_config.get("base_url") - self.api_key = embeddings_config.get("api_key") - self.max_model_len = embeddings_config.get("max_model_len", 8192) + 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 diff --git a/openrag/components/indexer/indexer.py b/openrag/components/indexer/indexer.py index 26508dbb2..e0f978776 100644 --- a/openrag/components/indexer/indexer.py +++ b/openrag/components/indexer/indexer.py @@ -17,20 +17,20 @@ config = load_config() save_uploaded_files = os.environ.get("SAVE_UPLOADED_FILES", "true").lower() == "true" -POOL_SIZE = config.ray.get("pool_size") -MAX_TASKS_PER_WORKER = config.ray.get("max_tasks_per_worker") +POOL_SIZE = config.ray.pool_size +MAX_TASKS_PER_WORKER = config.ray.max_tasks_per_worker @ray.remote( max_concurrency=config.ray.indexer.concurrency_groups.default, max_task_retries=config.ray.indexer.max_task_retries, concurrency_groups={ - "update": config.ray.indexer.concurrency_groups["update"], - "search": config.ray.indexer.concurrency_groups["search"], - "delete": config.ray.indexer.concurrency_groups["delete"], - "insert": config.ray.indexer.concurrency_groups["insert"], - "chunk": config.ray.indexer.concurrency_groups["chunk"], - "serialize": config.ray.indexer.concurrency_groups["serialize"], + "update": config.ray.indexer.concurrency_groups.update, + "search": config.ray.indexer.concurrency_groups.search, + "delete": config.ray.indexer.concurrency_groups.delete, + "insert": config.ray.indexer.concurrency_groups.insert, + "chunk": config.ray.indexer.concurrency_groups.chunk, + "serialize": config.ray.indexer.concurrency_groups.serialize, }, ) class Indexer: @@ -44,7 +44,7 @@ def __init__(self): self.chunker: BaseChunker = ChunkerFactory.create_chunker(self.config) self.default_partition = "_default" - self.enable_insertion = self.config.vectordb["enable"] + self.enable_insertion = self.config.vectordb.enable self.handle = ray.get_actor("Indexer", namespace="openrag") self.logger.info("Indexer actor initialized.") diff --git a/openrag/components/indexer/loaders/__init__.py b/openrag/components/indexer/loaders/__init__.py index 7b142edce..43fdafa90 100644 --- a/openrag/components/indexer/loaders/__init__.py +++ b/openrag/components/indexer/loaders/__init__.py @@ -15,7 +15,7 @@ logger = get_logger() -def get_loader_classes(config: dict) -> dict[str, type[BaseLoader]]: +def get_loader_classes(config) -> dict[str, type[BaseLoader]]: # 1. Discover all subclasses root_pkg = "components.indexer.loaders" root_path = Path(__file__).parent @@ -37,7 +37,7 @@ def get_loader_classes(config: dict) -> dict[str, type[BaseLoader]]: # 2. Read your config map of extensions β†’ class names loader_classes: dict[str, type[BaseLoader]] = {} - file_loaders = config.get("loader", {}).get("file_loaders", {}) + file_loaders = config.loader.file_loaders.model_dump() for ext, cls_name in file_loaders.items(): cls = discovered.get(cls_name) diff --git a/openrag/components/indexer/loaders/audio/local_whisper.py b/openrag/components/indexer/loaders/audio/local_whisper.py index 6cd594388..244682d61 100644 --- a/openrag/components/indexer/loaders/audio/local_whisper.py +++ b/openrag/components/indexer/loaders/audio/local_whisper.py @@ -15,11 +15,11 @@ if torch.cuda.is_available(): - WHISPER_NUM_GPUS = config.loader.local_whisper.get("whisper_num_gpus", 0.01) + WHISPER_NUM_GPUS = config.loader.local_whisper.whisper_num_gpus else: # On CPU WHISPER_NUM_GPUS = 0 -WHISPER_CONCURRENCY_PER_WORKER = config.loader.local_whisper.get("whisper_concurency_per_worker", 2) +WHISPER_CONCURRENCY_PER_WORKER = config.loader.local_whisper.whisper_concurrency_per_worker @ray.remote( @@ -36,7 +36,7 @@ def __init__(self): device = "cuda" if torch.cuda.is_available() else "cpu" compute_type = "float16" if device == "cuda" else "int8" - model_name = self.config.loader.local_whisper.get("model", "base") + model_name = self.config.loader.local_whisper.model self.logger.info("Loading Whisper model", model_name=model_name, device=device, compute_type=compute_type) self.model = WhisperModel(model_name, device=device, compute_type=compute_type) @@ -74,7 +74,7 @@ def __init__(self): self.logger = get_logger() - n_workers = config.loader.local_whisper.get("whisper_n_workers") + n_workers = config.loader.local_whisper.whisper_n_workers self.logger.info(f"Starting WhisperPool with {n_workers} workers") self.workers = [WhisperActor.remote() for _ in range(n_workers)] self._pending = [0] * n_workers diff --git a/openrag/components/indexer/loaders/base.py b/openrag/components/indexer/loaders/base.py index 288ff8272..5a5582504 100644 --- a/openrag/components/indexer/loaders/base.py +++ b/openrag/components/indexer/loaders/base.py @@ -37,7 +37,7 @@ class BaseLoader(ABC): def __init__(self, **kwargs) -> None: self.page_sep = "[PAGE_SEP]" self.config = kwargs.get("config") - settings: dict = dict(self.config.vlm) + settings: dict = self.config.vlm.model_dump() model_settings = { "temperature": 0.2, "max_retries": 3, @@ -46,8 +46,8 @@ def __init__(self, **kwargs) -> None: } settings.update(model_settings) - self.image_captioning = self.config.loader.get("image_captioning", True) - self.image_captioning_url = self.config.loader.get("image_captioning_url", True) + self.image_captioning = self.config.loader.image_captioning + self.image_captioning_url = self.config.loader.image_captioning_url self.vlm_endpoint = ChatOpenAI(**settings).with_retry(stop_after_attempt=2) diff --git a/openrag/components/indexer/loaders/eml_loader.py b/openrag/components/indexer/loaders/eml_loader.py index c23bcca72..96a5ee1e2 100644 --- a/openrag/components/indexer/loaders/eml_loader.py +++ b/openrag/components/indexer/loaders/eml_loader.py @@ -188,7 +188,7 @@ async def aload_document(self, file_path, metadata: dict | None = None, save_mar # Try fallback processing for images if file_ext in [".png", ".jpg", ".jpeg", ".svg"]: try: - if self.config.loader.get("image_captioning", False): + if self.image_captioning: # Try to load image directly from bytes as fallback image = Image.open(io.BytesIO(attachment["raw"])) caption = await self.get_image_description(image_data=image) @@ -221,12 +221,16 @@ async def aload_document(self, file_path, metadata: dict | None = None, save_mar os.unlink(temp_file_path) # Special handling for images with captioning if no specific loader or captioning is enabled - elif file_ext in [ - ".png", - ".jpg", - ".jpeg", - ".svg", - ] and self.config.loader.get("image_captioning", False): + elif ( + file_ext + in [ + ".png", + ".jpg", + ".jpeg", + ".svg", + ] + and self.image_captioning + ): try: # Load image from raw bytes image = Image.open(io.BytesIO(attachment["raw"])) diff --git a/openrag/components/indexer/loaders/pdf_loaders/docling2.py b/openrag/components/indexer/loaders/pdf_loaders/docling2.py index 4c19de9ed..238f1bd06 100644 --- a/openrag/components/indexer/loaders/pdf_loaders/docling2.py +++ b/openrag/components/indexer/loaders/pdf_loaders/docling2.py @@ -26,11 +26,11 @@ if torch.cuda.is_available(): - DOCLING_NUM_GPUS = config.loader.get("docling_num_gpus", 0.01) + DOCLING_NUM_GPUS = config.loader.docling_num_gpus else: # On CPU DOCLING_NUM_GPUS = 0 -DOCLING_MAX_TASKS_PER_WORKER = config.loader.get("docling_max_tasks_per_worker", 2) +DOCLING_MAX_TASKS_PER_WORKER = config.loader.docling_max_tasks_per_worker @ray.remote(num_gpus=DOCLING_NUM_GPUS) @@ -70,7 +70,7 @@ def __init__(self): self.logger = get_logger() self.config = load_config() - self.pool_size = config.loader.get("docling_pool_size", 1) + self.pool_size = self.config.loader.docling_pool_size self.actors = [DoclingWorker.remote() for _ in range(self.pool_size)] self._queue: asyncio.Queue[ray.actor.ActorHandle] = asyncio.Queue() @@ -109,7 +109,7 @@ async def aload_document(self, file_path, metadata, save_markdown=False): s += f"\n[PAGE_{i}]\n" enriched_content = s - if self.config.loader["image_captioning"]: + if self.image_captioning: pictures = result.document.pictures descriptions = await self.get_captions(pictures) for description in descriptions: diff --git a/openrag/components/indexer/loaders/pdf_loaders/marker.py b/openrag/components/indexer/loaders/pdf_loaders/marker.py index 59f1d1560..a72ba55c7 100644 --- a/openrag/components/indexer/loaders/pdf_loaders/marker.py +++ b/openrag/components/indexer/loaders/pdf_loaders/marker.py @@ -17,7 +17,7 @@ config = load_config() if torch.cuda.is_available(): - MARKER_NUM_GPUS = config.loader.get("marker_num_gpus", 0.01) + MARKER_NUM_GPUS = config.loader.marker_num_gpus else: # On CPU MARKER_NUM_GPUS = 0 @@ -34,13 +34,13 @@ def __init__(self): self.config = load_config() self.page_sep = "[PAGE_SEP]" - self._workers = self.config.loader.get("marker_max_processes") + self._workers = self.config.loader.marker_max_processes self.converter_config = { "output_format": "markdown", "paginate_output": True, "page_separator": self.page_sep, - "pdftext_workers": self.config.loader.get("marker_pdftext_workers"), + "pdftext_workers": self.config.loader.marker_pdftext_workers, "disable_multiprocessing": False, } os.environ["RAY_ADDRESS"] = "auto" @@ -89,7 +89,7 @@ def setup_mp(self): initializer=self._worker_init, initargs=(self.model_dict,), mp_context=mp.get_context("spawn"), - max_tasks_per_child=self.config.loader.get("marker_max_tasks_per_child", 5), + max_tasks_per_child=self.config.loader.marker_max_tasks_per_child, ) self.logger.info("MarkerWorker initialized with ProcessPoolExecutor") @@ -125,7 +125,7 @@ async def process_pdf(self, file_path: str): converter_config = self.converter_config.copy() loop = asyncio.get_event_loop() - timeout = self.config.loader.get("marker_timeout", 3600) + timeout = self.config.loader.marker_timeout def run_with_timeout(): future = self.executor.submit(self._process_pdf, file_path, converter_config) @@ -166,9 +166,9 @@ def __init__(self): self.logger = get_logger() self.config = load_config() - self.min_processes = self.config.loader.get("marker_min_processes") - self.max_processes = self.config.loader.get("marker_max_processes") - self.pool_size = config.loader.get("marker_pool_size") + self.min_processes = self.config.loader.marker_min_processes + self.max_processes = self.config.loader.marker_max_processes + self.pool_size = self.config.loader.marker_pool_size self.actors = [MarkerWorker.remote() for _ in range(self.pool_size)] self._queue: asyncio.Queue[ray.actor.ActorHandle] = asyncio.Queue() @@ -199,7 +199,7 @@ async def process_pdf(self, file_path: str): # Ensure the worker pool is healthy await self.ensure_worker_pool_healthy(worker) try: - timeout = self.config.loader.get("marker_timeout", 3600) + timeout = self.config.loader.marker_timeout future = worker.process_pdf.remote(file_path) return await call_ray_actor_with_timeout( future, @@ -232,7 +232,7 @@ async def aload_document( start = time.time() try: - timeout = self.config.loader.get("marker_timeout", 3600) + timeout = self.config.loader.marker_timeout future = self.worker.process_pdf.remote(file_path_str) markdown, images = await call_ray_actor_with_timeout( future, diff --git a/openrag/components/indexer/loaders/pdf_loaders/openai.py b/openrag/components/indexer/loaders/pdf_loaders/openai.py index 06535ae57..9aec8ba24 100644 --- a/openrag/components/indexer/loaders/pdf_loaders/openai.py +++ b/openrag/components/indexer/loaders/pdf_loaders/openai.py @@ -29,15 +29,15 @@ def __init__(self, **kwargs): super().__init__(**kwargs) self.llm = ChatOpenAI( - base_url=self.config.loader["openai"]["base_url"], - api_key=self.config.loader["openai"]["api_key"], - model=self.config.loader["openai"]["model"], - temperature=self.config.loader["openai"].get("temperature", 0.2), - timeout=self.config.loader["openai"].get("timeout", 180), - max_retries=self.config.loader["openai"].get("max_retries", 2), - top_p=self.config.loader["openai"].get("top_p", 0.9), + base_url=self.config.loader.openai.base_url, + api_key=self.config.loader.openai.api_key, + model=self.config.loader.openai.model, + temperature=self.config.loader.openai.temperature, + timeout=self.config.loader.openai.timeout, + max_retries=self.config.loader.openai.max_retries, + top_p=self.config.loader.openai.top_p, ) - self.llm_semaphore = asyncio.Semaphore(self.config.loader["openai"].get("concurrency_limit", 20)) + self.llm_semaphore = asyncio.Semaphore(self.config.loader.openai.concurrency_limit) async def aload_document( self, @@ -77,7 +77,7 @@ async def _assemble_markdown(self, pages: list[Image.Image], results: list[dict] for page_img, page_res in zip(pages, results): if not page_res: continue - if self.config["loader"]["image_captioning"]: + if self.image_captioning: await self._caption_images(page_img, page_res) markdown_parts.append(self._result_to_md(page_res)) return "\n\n".join(markdown_parts).strip() diff --git a/openrag/components/indexer/loaders/serializer.py b/openrag/components/indexer/loaders/serializer.py index 5ceccbef4..28e278892 100644 --- a/openrag/components/indexer/loaders/serializer.py +++ b/openrag/components/indexer/loaders/serializer.py @@ -12,11 +12,11 @@ # Set ray resources if torch.cuda.is_available(): - NUM_GPUS = config.ray.get("num_gpus") + NUM_GPUS = config.ray.num_gpus else: # On CPU NUM_GPUS = 0 -DICT_MIMETYPES = dict(config.loader["mimetypes"]) +DICT_MIMETYPES = config.loader.mimetypes.to_dict() @ray.remote(max_restarts=5) @@ -30,7 +30,7 @@ def __init__(self, data_dir=None, **kwargs) -> None: self.data_dir = data_dir self.kwargs = kwargs self.kwargs["config"] = self.config - self.save_markdown = self.config.loader.get("save_markdown", False) + self.save_markdown = self.config.loader.save_markdown # Initialize loader classes: self.loader_classes = get_loader_classes(config=self.config) diff --git a/openrag/components/indexer/loaders/test_doc_loader.py b/openrag/components/indexer/loaders/test_doc_loader.py index 1d2a32acc..68d2f2b30 100644 --- a/openrag/components/indexer/loaders/test_doc_loader.py +++ b/openrag/components/indexer/loaders/test_doc_loader.py @@ -9,6 +9,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest +from config.models import LoaderConfig, VLMConfig from langchain_core.documents.base import Document as LCDocument @@ -16,8 +17,8 @@ def mock_config(): """Create a minimal mock config for BaseLoader.""" config = MagicMock() - config.vlm = {"model": "mock", "base_url": "http://mock", "api_key": "mock"} - config.loader = {"image_captioning": False, "image_captioning_url": False} + config.vlm = VLMConfig(model="mock", base_url="http://mock", api_key="mock") + config.loader = LoaderConfig(image_captioning=False, image_captioning_url=False) return config diff --git a/openrag/components/indexer/utils/files.py b/openrag/components/indexer/utils/files.py index f9194de46..f3d11b48f 100644 --- a/openrag/components/indexer/utils/files.py +++ b/openrag/components/indexer/utils/files.py @@ -9,7 +9,7 @@ from fastapi import UploadFile config = load_config() -SERIALIZE_TIMEOUT = config.ray.indexer.get("serialize_timeout", 3600) +SERIALIZE_TIMEOUT = config.ray.indexer.serialize_timeout def sanitize_filename(filename: str) -> str: diff --git a/openrag/components/indexer/vectordb/utils.py b/openrag/components/indexer/vectordb/utils.py index f079f0d29..5c60c9838 100644 --- a/openrag/components/indexer/vectordb/utils.py +++ b/openrag/components/indexer/vectordb/utils.py @@ -19,7 +19,7 @@ logger = get_logger() config = load_config() -DEFAULT_FILE_QUOTA = config.rdb.get("default_file_quota", -1) +DEFAULT_FILE_QUOTA = config.rdb.default_file_quota class PartitionFileManager: diff --git a/openrag/components/indexer/vectordb/vectordb.py b/openrag/components/indexer/vectordb/vectordb.py index 0c296dc67..16b26e894 100644 --- a/openrag/components/indexer/vectordb/vectordb.py +++ b/openrag/components/indexer/vectordb/vectordb.py @@ -145,8 +145,8 @@ def __init__(self): self.logger = get_logger() # init milvus clients - self.port = self.config.vectordb.get("port") - self.host = self.config.vectordb.get("host") + self.port = self.config.vectordb.port + self.host = self.config.vectordb.host uri = f"http://{self.host}:{self.port}" self.uri = uri try: @@ -162,7 +162,7 @@ def __init__(self): # embedder self.embedder: BaseEmbedding = EmbeddingFactory.get_embedder(embeddings_config=self.config.embedder) - self.hybrid_search = self.config.vectordb.get("hybrid_search", True) + self.hybrid_search = self.config.vectordb.hybrid_search # partition related params self.rdb_host = self.config.rdb.host self.rdb_port = self.config.rdb.port @@ -171,7 +171,7 @@ def __init__(self): self.partition_file_manager: PartitionFileManager = None # Initialize collection-related attributes - self.collection_name = self.config.vectordb.get("collection_name", "vdb_test") + self.collection_name = self.config.vectordb.collection_name self.collection_loaded = False self.load_collection() @@ -1208,7 +1208,7 @@ class ConnectorFactory: @staticmethod def get_vectordb_cls(): - name = config.vectordb.get("connector_name") + name = config.vectordb.connector_name vdb_cls = ConnectorFactory.CONNECTORS.get(name) if not vdb_cls: raise ValueError(f"VECTORDB '{name}' is not supported.") diff --git a/openrag/components/llm.py b/openrag/components/llm.py index 730be1b58..bfaf4ed76 100644 --- a/openrag/components/llm.py +++ b/openrag/components/llm.py @@ -10,7 +10,7 @@ class LLM: def __init__(self, llm_config, logger=None): self.logger = logger - default_llm_config = dict(llm_config) + 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 diff --git a/openrag/components/map_reduce.py b/openrag/components/map_reduce.py index e9f89f348..af2d4ce47 100644 --- a/openrag/components/map_reduce.py +++ b/openrag/components/map_reduce.py @@ -49,13 +49,13 @@ class SummarizedChunk(BaseModel): class RAGMapReduce: def __init__(self, config): self.config = config - self.slm: ChatOpenAI = ChatOpenAI(**config.llm).with_structured_output(SummarizedChunk) + self.slm: ChatOpenAI = ChatOpenAI(**config.llm.model_dump()).with_structured_output(SummarizedChunk) map_reduce_config = self.config.map_reduce - self.initial_batch_size = map_reduce_config["initial_batch_size"] - self.expansion_batch_size = map_reduce_config["expansion_batch_size"] - self.max_total_documents = map_reduce_config["max_total_documents"] + self.initial_batch_size = map_reduce_config.initial_batch_size + self.expansion_batch_size = map_reduce_config.expansion_batch_size + self.max_total_documents = map_reduce_config.max_total_documents - self.debug = map_reduce_config.get("debug", True) + self.debug = map_reduce_config.debug assert self.max_total_documents >= self.initial_batch_size, ( "`max_total_documents` must be greater than or equal to `initial_batch_size`" diff --git a/openrag/components/pipeline.py b/openrag/components/pipeline.py index fffd433c6..67fec45c6 100644 --- a/openrag/components/pipeline.py +++ b/openrag/components/pipeline.py @@ -26,7 +26,7 @@ logger = get_logger() config = load_config() -VECTORDB_TIMEOUT = config.ray.indexer.get("vectordb_timeout", 30) +VECTORDB_TIMEOUT = config.ray.indexer.vectordb_timeout class RAGMODE(Enum): @@ -49,10 +49,10 @@ def __init__(self) -> None: self.retriever: BaseRetriever = RetrieverFactory.create_retriever(config=config) # reranker - self.reranker_enabled = config.reranker["enable"] + self.reranker_enabled = config.reranker.enable self.reranker = Reranker(logger, config) logger.debug("Reranker", enabled=self.reranker_enabled) - self.reranker_top_k = config.reranker["top_k"] + self.reranker_top_k = config.reranker.top_k async def retrieve_docs( self, @@ -120,19 +120,19 @@ def __init__(self) -> None: self.retriever_pipeline = RetrieverPipeline() # RAG - self.rag_mode = config.rag["mode"] - self.chat_history_depth = config.rag["chat_history_depth"] - self.max_context_tokens = config.reranker.get("top_k", 10) * config.chunker.get("chunk_size", 512) + self.rag_mode = config.rag.mode + self.chat_history_depth = config.rag.chat_history_depth + self.max_context_tokens = config.reranker.top_k * config.chunker.chunk_size self.llm_client = LLM(config.llm, logger) self.query_generator = ChatOpenAI( - base_url=config.llm.get("base_url"), - api_key=config.llm.get("api_key"), - model=config.llm.get("model"), - temperature=config.llm.get("temperature", 0.3), + base_url=config.llm.base_url, + api_key=config.llm.api_key, + model=config.llm.model, + temperature=config.llm.temperature, ).with_structured_output(SearchQueries, method="function_calling") - self.max_contextualized_query_len = config.rag["max_contextualized_query_len"] + self.max_contextualized_query_len = config.rag.max_contextualized_query_len # map reduce self.map_reduce: RAGMapReduce = RAGMapReduce(config=config) @@ -140,7 +140,7 @@ def __init__(self) -> None: # Web search self.web_search_service = WebSearchFactory.create_service(config) if self.web_search_service.provider: - logger.info("Web search enabled", provider=config.websearch.get("provider")) + logger.info("Web search enabled", provider=config.websearch.provider) else: logger.info("Web search disabled (WEBSEARCH_API_TOKEN not set)") @@ -207,7 +207,7 @@ async def _prepare_for_chat_completion(self, partition: list[str] | None, payloa ) # 2. get docs and/or web results concurrently - top_k = config.map_reduce["max_total_documents"] if use_map_reduce else None + top_k = config.map_reduce.max_total_documents if use_map_reduce else None if workspace: vectordb = ray.get_actor("Vectordb", namespace="openrag") ws = await call_ray_actor_with_timeout( diff --git a/openrag/components/prompts/prompts.py b/openrag/components/prompts/prompts.py index e7cf0ec6d..afb2c3727 100644 --- a/openrag/components/prompts/prompts.py +++ b/openrag/components/prompts/prompts.py @@ -5,15 +5,15 @@ config = load_config() prompts_dir: Path = config.paths.prompts_dir -prompt_mapping: dict = config.prompts +prompt_mapping = config.prompts def load_prompt( prompt_name: str, prompts_dir: Path = prompts_dir, - prompt_mapping: dict = prompt_mapping, + prompt_mapping=prompt_mapping, ) -> str: - file_name = prompt_mapping.get(prompt_name, None) + file_name = getattr(prompt_mapping, prompt_name, None) if not file_name: raise ValueError(f"No associated file name found for prompt: `{prompt_name}`") diff --git a/openrag/components/reranker.py b/openrag/components/reranker.py index a40e3ef71..14403732a 100644 --- a/openrag/components/reranker.py +++ b/openrag/components/reranker.py @@ -42,8 +42,8 @@ def rrf_reranking(doc_lists: list[list], k: int = 60) -> list[Document]: 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.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) diff --git a/openrag/components/retriever.py b/openrag/components/retriever.py index 5c0c87f58..3f074adaf 100644 --- a/openrag/components/retriever.py +++ b/openrag/components/retriever.py @@ -9,7 +9,6 @@ from langchain_core.output_parsers import StrOutputParser from langchain_core.prompts import ChatPromptTemplate from langchain_openai import ChatOpenAI -from omegaconf import OmegaConf from utils.dependencies import get_vectordb from utils.logger import get_logger @@ -321,14 +320,14 @@ class RetrieverFactory: } @classmethod - def create_retriever(cls, config: OmegaConf) -> ABCRetriever: - retreiverConfig = OmegaConf.to_container(config.retriever, resolve=True) + def create_retriever(cls, config) -> ABCRetriever: + retrieverConfig = config.retriever.model_dump() - retriever_type = retreiverConfig.pop("type") + retriever_type = retrieverConfig.pop("type") retriever_cls = RetrieverFactory.RETRIEVERS.get(retriever_type, None) if retriever_cls is None: raise ValueError(f"Unknown retriever type: {retriever_type}") - retreiverConfig["llm"] = ChatOpenAI(**config.llm) - return retriever_cls(**retreiverConfig) + retrieverConfig["llm"] = ChatOpenAI(**config.llm.model_dump()) + return retriever_cls(**retrieverConfig) diff --git a/openrag/components/test_llm.py b/openrag/components/test_llm.py index 122d57b64..822313be3 100644 --- a/openrag/components/test_llm.py +++ b/openrag/components/test_llm.py @@ -1,16 +1,17 @@ import pytest from components.llm import LLM +from config.models import LLMConfig @pytest.fixture def llm(): return LLM( - { - "base_url": "http://default-llm:8000/v1", - "api_key": "default-key", - "model": "default-model", - "temperature": 0.3, - } + LLMConfig( + base_url="http://default-llm:8000/v1", + api_key="default-key", + model="default-model", + temperature=0.3, + ) ) diff --git a/openrag/components/utils.py b/openrag/components/utils.py index cd871ae24..3e3055508 100644 --- a/openrag/components/utils.py +++ b/openrag/components/utils.py @@ -90,7 +90,7 @@ async def __aexit__(self, exc_type, exc, tb): def get_num_tokens(): global _cached_length_function if _cached_length_function is None: - llm = ChatOpenAI(**config.llm) + llm = ChatOpenAI(**config.llm.model_dump()) _cached_length_function = llm.get_num_tokens return _cached_length_function diff --git a/openrag/components/websearch/__init__.py b/openrag/components/websearch/__init__.py index 45df91ac5..912d328d3 100644 --- a/openrag/components/websearch/__init__.py +++ b/openrag/components/websearch/__init__.py @@ -12,33 +12,33 @@ class WebSearchFactory: @staticmethod def create_service(config) -> WebSearchService: - """Create a WebSearchService from Hydra config, following the embedder/retriever pattern.""" + """Create a WebSearchService from config, following the embedder/retriever pattern.""" ws_config = config.websearch - api_token = ws_config.get("api_token", "") - max_tokens = ws_config.get("max_tokens") + api_token = ws_config.api_token + max_tokens = ws_config.max_tokens if not api_token: return WebSearchService(provider=None, max_tokens=max_tokens) - provider_name = ws_config.get("provider", "") + provider_name = ws_config.provider provider_cls = PROVIDER_MAPPING.get(provider_name) if provider_cls is None: raise ValueError(f"Unsupported web search provider: {provider_name}") provider = provider_cls( api_token=api_token, - base_url=ws_config.get("base_url"), - top_k=ws_config.get("top_k"), - lang=ws_config.get("lang"), + base_url=ws_config.base_url, + top_k=ws_config.top_k, + lang=ws_config.lang, ) content_fetcher = None - if ws_config.get("fetch_content"): + if ws_config.fetch_content: content_fetcher = ContentFetcher( - max_results=ws_config.get("fetch_max_results"), - timeout=ws_config.get("fetch_timeout"), - max_tokens_per_page=ws_config.get("fetch_max_tokens"), - verify_ssl=ws_config.get("fetch_verify_ssl"), + max_results=ws_config.fetch_max_results, + timeout=ws_config.fetch_timeout, + max_tokens_per_page=ws_config.fetch_max_tokens, + verify_ssl=ws_config.fetch_verify_ssl, ) return WebSearchService( diff --git a/openrag/config/__init__.py b/openrag/config/__init__.py index 3827607ac..a2e454a45 100644 --- a/openrag/config/__init__.py +++ b/openrag/config/__init__.py @@ -1,3 +1,37 @@ -from .config import load_config +"""OpenRAG configuration package. -__all__ = [load_config] +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 .models 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/config/config.py b/openrag/config/config.py deleted file mode 100644 index f7d9d6c13..000000000 --- a/openrag/config/config.py +++ /dev/null @@ -1,27 +0,0 @@ -import os -from pathlib import Path - -from dotenv import load_dotenv -from hydra import compose, initialize_config_dir -from hydra.core.global_hydra import GlobalHydra -from omegaconf import OmegaConf - -CONFIG_PATH = Path(os.environ.get("CONFIG_PATH", "/app/.hydra_config")).resolve() - - -def load_config(config_path=CONFIG_PATH, overrides=None) -> OmegaConf: - load_dotenv() - - # Clear existing Hydra instance to prevent "already initialized" errors - if GlobalHydra.instance().is_initialized(): - GlobalHydra.instance().clear() - - # TODO: I set the version base to 1.1 to silence the warning message, review how we want to handle versioning - with initialize_config_dir(config_dir=str(config_path), job_name="config_loader", version_base="1.1"): - config = compose(config_name="config", overrides=overrides) - - config.paths.data_dir = Path(config.paths.data_dir).resolve() - config.paths.log_dir = Path(config.paths.log_dir).resolve() - config.paths.prompts_dir = Path(config.paths.prompts_dir).resolve() - - return config diff --git a/openrag/config/loader.py b/openrag/config/loader.py new file mode 100644 index 000000000..60b6f4a8f --- /dev/null +++ b/openrag/config/loader.py @@ -0,0 +1,291 @@ +"""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 .models 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), + # 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_ENABLED", "reranker.enable", bool), + ("RERANKER_MODEL", "reranker.model_name", str), + ("RERANKER_TOP_K", "reranker.top_k", int), + ("RERANKER_BASE_URL", "reranker.base_url", str), + # 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_MIN_PROCESSES", "loader.marker_min_processes", int), + ("MARKER_NUM_GPUS", "loader.marker_num_gpus", float), + ("MARKER_TIMEOUT", "loader.marker_timeout", int), + ("MARKER_PDFTEXT_WORKERS", "loader.marker_pdftext_workers", 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), + ("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), + # 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 + + # RERANKER_PORT: build default base_url if RERANKER_BASE_URL not set + if not os.environ.get("RERANKER_BASE_URL"): + port = os.environ.get("RERANKER_PORT", "7997") + reranker = data.setdefault("reranker", {}) + if not reranker.get("base_url"): + reranker["base_url"] = f"http://reranker:{port}" + + 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) + + # 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 new file mode 100644 index 000000000..50564afce --- /dev/null +++ b/openrag/config/models.py @@ -0,0 +1,477 @@ +"""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 + + +# --------------------------------------------------------------------------- +# 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 + + +# --------------------------------------------------------------------------- +# 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 RerankerConfig(ConfigMixin): + enable: bool = True + model_name: str = "Alibaba-NLP/gte-multilingual-reranker-base" + top_k: int = 10 + base_url: str = "" + + +# --------------------------------------------------------------------------- +# 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) +# --------------------------------------------------------------------------- +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 + + +# --------------------------------------------------------------------------- +# 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 + + +# --------------------------------------------------------------------------- +# 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 = 10 + marker_pool_size: int = 1 + marker_max_processes: int = 2 + marker_min_processes: int = 1 + marker_num_gpus: float = 0.01 + marker_timeout: int = 3600 + marker_pdftext_workers: int = 2 + 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) + transcriber: TranscriberConfig = Field(default_factory=TranscriberConfig) + openai: OpenAILoaderConfig = Field(default_factory=OpenAILoaderConfig) + + +# --------------------------------------------------------------------------- +# 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 + + +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 = False + + +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=RerankerConfig) + 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/models/openai.py b/openrag/models/openai.py index 323e44d64..b77305f70 100644 --- a/openrag/models/openai.py +++ b/openrag/models/openai.py @@ -4,7 +4,7 @@ from pydantic import BaseModel, Field config = load_config() -default_max_tokens = int(config.llm_context.get("max_output_tokens", 1024)) +default_max_tokens = config.llm_context.max_output_tokens # Classes pour la compatibilitΓ© OpenAI diff --git a/openrag/routers/indexer.py b/openrag/routers/indexer.py index 367f34dac..861a75080 100644 --- a/openrag/routers/indexer.py +++ b/openrag/routers/indexer.py @@ -39,15 +39,14 @@ # load config config = load_config() DATA_DIR = config.paths.data_dir -VECTORDB_TIMEOUT = config.ray.indexer.get("vectordb_timeout", 30) +VECTORDB_TIMEOUT = config.ray.indexer.vectordb_timeout FORBIDDEN_CHARS_IN_FILE_ID = set("/") # set('"<>#%{}|\\^`[]') LOG_FILE = Path(config.paths.log_dir or "logs") / "app.json" -VECTORDB_TIMEOUT = config.ray.indexer.get("vectordb_timeout", 30) # supported file formats or mimetypes -ACCEPTED_FILE_FORMATS = dict(config.loader["file_loaders"]).keys() -DICT_MIMETYPES = dict(config.loader["mimetypes"]) +ACCEPTED_FILE_FORMATS = config.loader.file_loaders.model_dump().keys() +DICT_MIMETYPES = config.loader.mimetypes.to_dict() # URL scheme configuration PREFERRED_URL_SCHEME = config.server.preferred_url_scheme diff --git a/openrag/routers/openai.py b/openrag/routers/openai.py index 9eb686fb7..e659c59e8 100644 --- a/openrag/routers/openai.py +++ b/openrag/routers/openai.py @@ -148,7 +148,7 @@ def is_direct_llm_model( Returns True if model is None, empty, or matches the configured default model. """ - return request.model is None or request.model == "" or request.model == config.llm.get("model") + return request.model is None or request.model == "" or request.model == config.llm.model async def _fetch_max_model_tokens() -> int: @@ -157,10 +157,10 @@ async def _fetch_max_model_tokens() -> int: Queries `/v1/models` and looks for `max_model_len` for the configured LLM model. Falls back to `config.llm_context.max_llm_context_size` (default 8192) if unavailable. """ - default_limit = int(config.llm_context.get("max_llm_context_size", 8192)) - model_id = config.llm.get("model") + default_limit = int(config.llm_context.max_llm_context_size) + model_id = config.llm.model try: - openai_models = await get_openai_models(base_url=config.llm["base_url"], api_key=config.llm["api_key"]) + openai_models = await get_openai_models(base_url=config.llm.base_url, api_key=config.llm.api_key) model = next((m for m in openai_models if m.id == model_id), None) if model is None: logger.warning(f"No model found for {model_id}. Using default context size.") @@ -185,7 +185,7 @@ def get_max_model_tokens() -> int: """Return the cached max model token limit (populated at startup).""" if _max_model_tokens is not None: return _max_model_tokens - return int(config.llm_context.get("max_llm_context_size", 8192)) + return int(config.llm_context.max_llm_context_size) def validate_tokens_limit( @@ -206,7 +206,7 @@ def validate_tokens_limit( if isinstance(request, OpenAIChatCompletionRequest): message_tokens = sum(_length_function(m.content or "") + 4 for m in request.messages) - default_output_tokens = int(config.llm_context.get("max_output_tokens", 1024)) + default_output_tokens = int(config.llm_context.max_output_tokens) requested_tokens = request.max_tokens or default_output_tokens total_tokens_needed = message_tokens + requested_tokens @@ -229,7 +229,7 @@ def validate_tokens_limit( elif isinstance(request, OpenAICompletionRequest): prompt_tokens = _length_function(request.prompt) - default_output_tokens = int(config.llm_context.get("max_output_tokens", 1024)) + default_output_tokens = int(config.llm_context.max_output_tokens) requested_tokens = request.max_tokens or default_output_tokens total_tokens_needed = prompt_tokens + requested_tokens @@ -310,7 +310,7 @@ async def openai_chat_completion( user_partitions=Depends(current_user_or_admin_partitions_list), _: None = Depends(check_llm_model_availability), ): - model_name = request.model or config.llm.get("model") + model_name = request.model or config.llm.model log = logger.bind(model=model_name, endpoint="/chat/completions") if not request.messages or request.messages[-1].role != "user" or not request.messages[-1].content: @@ -404,7 +404,7 @@ async def openai_completion( user_partitions=Depends(current_user_or_admin_partitions_list), _: None = Depends(check_llm_model_availability), ): - model_name = request.model or config.llm.get("model") + model_name = request.model or config.llm.model log = logger.bind(model=model_name, endpoint="/completions") if not request.prompt: diff --git a/openrag/routers/search.py b/openrag/routers/search.py index 888699246..6b19e41a1 100644 --- a/openrag/routers/search.py +++ b/openrag/routers/search.py @@ -15,7 +15,7 @@ ) _config = load_config() -VECTORDB_TIMEOUT = _config.ray.indexer.get("vectordb_timeout", 30) +VECTORDB_TIMEOUT = _config.ray.indexer.vectordb_timeout logger = get_logger() diff --git a/openrag/routers/utils.py b/openrag/routers/utils.py index 0336cdfc6..1e1d3e6e6 100644 --- a/openrag/routers/utils.py +++ b/openrag/routers/utils.py @@ -23,8 +23,8 @@ LOG_FILE = Path(config.paths.log_dir or "logs") / "app.json" # supported file formats or mimetypes -ACCEPTED_FILE_FORMATS = dict(config.loader["file_loaders"]).keys() -DICT_MIMETYPES = dict(config.loader["mimetypes"]) +ACCEPTED_FILE_FORMATS = config.loader.file_loaders.model_dump().keys() +DICT_MIMETYPES = config.loader.mimetypes.to_dict() ROLE_HIERARCHY = { "viewer": 1, @@ -33,7 +33,7 @@ } # File quota per user -DEFAULT_FILE_QUOTA = config.rdb.get("default_file_quota", -1) +DEFAULT_FILE_QUOTA = config.rdb.default_file_quota def current_user(request: Request): @@ -307,9 +307,9 @@ async def get_openai_models(base_url: str, api_key: str, timeout: int = 30): async def check_llm_model_availability(request: Request): llm_param = config.llm - base_url = llm_param.get("base_url") - model = llm_param.get("model") - api_key = llm_param.get("api_key") + base_url = llm_param.base_url + model = llm_param.model + api_key = llm_param.api_key missing = [k for k, v in {"base_url": base_url, "model": model, "api_key": api_key}.items() if not v] if missing: @@ -323,7 +323,7 @@ async def check_llm_model_availability(request: Request): try: log.debug("Validating model") - timeout = int(llm_param.get("timeout", 30)) + timeout = int(llm_param.timeout) openai_models = await get_openai_models(base_url=base_url, api_key=api_key, timeout=timeout) available_models = {m.id for m in openai_models} if model not in available_models: diff --git a/openrag/routers/workspaces.py b/openrag/routers/workspaces.py index fefc15331..1d3a9663f 100644 --- a/openrag/routers/workspaces.py +++ b/openrag/routers/workspaces.py @@ -16,7 +16,7 @@ logger = get_logger() _config = load_config() -VECTORDB_TIMEOUT = _config.ray.indexer.get("vectordb_timeout", 30) +VECTORDB_TIMEOUT = _config.ray.indexer.vectordb_timeout _WORKSPACE_ID_RE = re.compile(r"[a-zA-Z0-9_-]+") diff --git a/openrag/scripts/backup.py b/openrag/scripts/backup.py index 0becc584b..657583de6 100644 --- a/openrag/scripts/backup.py +++ b/openrag/scripts/backup.py @@ -186,9 +186,7 @@ def load_openrag_config(logger): logger: Logger instance. Returns: - tuple: - rdb (dict): Relational database configuration. - vdb (dict): Vector database configuration. + tuple: (RDBConfig, VectorDBConfig) Pydantic config models. """ from config import load_config @@ -198,7 +196,7 @@ def load_openrag_config(logger): logger.error(f"Failed while trying to obtain OpenRAG config: {e}") raise - return config["rdb"], config["vectordb"] + return config.rdb, config.vectordb # Arguments and configs import argparse @@ -222,20 +220,18 @@ def load_openrag_config(logger): rdb, vdb = load_openrag_config(logger) if args.verbose: - logger.info( - f"rdb @ {rdb['host']}:{rdb['port']} | vdb @ {vdb['host']}:{vdb['port']} | collection: {vdb['collection_name']}" - ) + logger.info(f"rdb @ {rdb.host}:{rdb.port} | vdb @ {vdb.host}:{vdb.port} | collection: {vdb.collection_name}") # List existing partitions try: pfm = PartitionFileManager( - database_url=f"postgresql://{rdb['user']}:{rdb['password']}@{rdb['host']}:{rdb['port']}/partitions_for_collection_{vdb['collection_name']}", + database_url=f"postgresql://{rdb.user}:{rdb.password}@{rdb.host}:{rdb.port}/partitions_for_collection_{vdb.collection_name}", logger=logger, ) existing_partitions = {item["partition"]: item for item in pfm.list_partitions()} except Exception as e: - logger.error(f"Failed while accessing PartitionFileManager at {rdb['host']}:{rdb['port']}\n{e}") + logger.error(f"Failed while accessing PartitionFileManager at {rdb.host}:{rdb.port}\n{e}") raise if args.include_only: @@ -259,15 +255,15 @@ def load_openrag_config(logger): # Connect to Milvus try: - connections.connect("default", host=vdb["host"], port=vdb["port"]) + connections.connect("default", host=vdb.host, port=vdb.port) except Exception as e: - logger.error(f"Can't connect to Milvus at {vdb['host']}:{vdb['port']}\n{e}") + logger.error(f"Can't connect to Milvus at {vdb.host}:{vdb.port}\n{e}") raise try: - vdb_collection = Collection(vdb["collection_name"]) + vdb_collection = Collection(vdb.collection_name) except Exception as e: - logger.error(f"Can't access Milvus collection {vdb['collection_name']} at {vdb['host']}:{vdb['port']}\n{e}") + logger.error(f"Can't access Milvus collection {vdb.collection_name} at {vdb.host}:{vdb.port}\n{e}") raise try: diff --git a/openrag/scripts/restore.py b/openrag/scripts/restore.py index 7986259bf..735aa6f8d 100644 --- a/openrag/scripts/restore.py +++ b/openrag/scripts/restore.py @@ -205,7 +205,7 @@ def main(): int: Exit code (0 on success, non-zero on failure). """ - def load_openrag_config(logger: Any) -> tuple[dict[str, Any], dict[str, Any]]: + def load_openrag_config(logger: Any): """ Loads OpenRAG configuration. @@ -213,9 +213,7 @@ def load_openrag_config(logger: Any) -> tuple[dict[str, Any], dict[str, Any]]: logger: Logger instance. Returns: - tuple: - rdb (dict): Relational database configuration. - vdb (dict): Vector database configuration. + tuple: (RDBConfig, VectorDBConfig) Pydantic config models. """ from config import load_config @@ -225,7 +223,7 @@ def load_openrag_config(logger: Any) -> tuple[dict[str, Any], dict[str, Any]]: logger.error(f"Failed while trying to obtain OpenRAG config: {e}") raise - return config["rdb"], config["vectordb"] + return config.rdb, config.vectordb # Arguments and configs import argparse @@ -269,20 +267,18 @@ def load_openrag_config(logger: Any) -> tuple[dict[str, Any], dict[str, Any]]: rdb, vdb = load_openrag_config(logger) if args.verbose: - logger.info( - f"rdb @ {rdb['host']}:{rdb['port']} | vdb @ {vdb['host']}:{vdb['port']} | collection: {vdb['collection_name']}" - ) + logger.info(f"rdb @ {rdb.host}:{rdb.port} | vdb @ {vdb.host}:{vdb.port} | collection: {vdb.collection_name}") # List existing partitions try: pfm = PartitionFileManager( - database_url=f"postgresql://{rdb['user']}:{rdb['password']}@{rdb['host']}:{rdb['port']}/partitions_for_collection_{vdb['collection_name']}", + database_url=f"postgresql://{rdb.user}:{rdb.password}@{rdb.host}:{rdb.port}/partitions_for_collection_{vdb.collection_name}", logger=logger, ) existing_partitions = {item["partition"]: item for item in pfm.list_partitions()} except Exception as e: - logger.error(f"Failed while accessing PartitionFileManager at {rdb['host']}:{rdb['port']}\n{e}") + logger.error(f"Failed while accessing PartitionFileManager at {rdb.host}:{rdb.port}\n{e}") raise if args.include_only: @@ -291,7 +287,7 @@ def load_openrag_config(logger: Any) -> tuple[dict[str, Any], dict[str, Any]]: logger.error(f'Partition "{part_name}" already exists') return 1 - client = MilvusClient(uri=f"http://{vdb['host']}:{vdb['port']}") + client = MilvusClient(uri=f"http://{vdb.host}:{vdb.port}") try: with open_backup_file(args.input, logger) as fh: @@ -316,7 +312,7 @@ def load_openrag_config(logger: Any) -> tuple[dict[str, Any], dict[str, Any]]: if line in ["vdb"]: read_vdb_section( fh, - vdb["collection_name"], + vdb.collection_name, added_documents, client, args.batch_size, diff --git a/openrag/tests/test_relationships_integration.py b/openrag/tests/test_relationships_integration.py index fbb53bfc4..7b984e3c0 100644 --- a/openrag/tests/test_relationships_integration.py +++ b/openrag/tests/test_relationships_integration.py @@ -291,7 +291,7 @@ class TestRelationshipAwareRetrieverIntegration: Tests the retriever's ability to expand search results with related and ancestor documents. - Note: These tests require the hydra configuration to be available. + Note: These tests require the configuration to be available. They are marked to skip when the config is not found. """ @@ -337,7 +337,7 @@ async def test_retriever_without_expansion_returns_base_results(self): try: from components.retriever import RelationshipAwareRetriever except Exception: - pytest.skip("Requires hydra config to be available") + pytest.skip("Requires config to be available") with patch("components.retriever.get_vectordb") as mock_get_db: mock_db = MagicMock() @@ -362,7 +362,7 @@ async def test_retriever_with_include_related_expands_results(self, mock_documen try: from components.retriever import RelationshipAwareRetriever except Exception: - pytest.skip("Requires hydra config to be available") + pytest.skip("Requires config to be available") with patch("components.retriever.get_vectordb") as mock_get_db: mock_db = MagicMock() @@ -394,7 +394,7 @@ async def test_retriever_deduplicates_results(self, mock_documents): try: from components.retriever import RelationshipAwareRetriever except Exception: - pytest.skip("Requires hydra config to be available") + pytest.skip("Requires config to be available") with patch("components.retriever.get_vectordb") as mock_get_db: mock_db = MagicMock() diff --git a/openrag/utils/dependencies.py b/openrag/utils/dependencies.py index d547925fc..ee21cd004 100644 --- a/openrag/utils/dependencies.py +++ b/openrag/utils/dependencies.py @@ -46,7 +46,7 @@ def get_serializer(): def get_marker_pool(): - pdf_loader = config.loader.file_loaders.get("pdf") + pdf_loader = config.loader.file_loaders.pdf match pdf_loader: case "DoclingLoader2": return get_or_create_actor("DoclingPool", DoclingPool, lifetime="detached") @@ -64,7 +64,7 @@ def get_vectordb(): def init_audio_actor(): - use_whisper_lang_detector = config.loader.transcriber.get("use_whisper_lang_detector", True) + use_whisper_lang_detector = config.loader.transcriber.use_whisper_lang_detector file_loaders = config.loader.file_loaders loader_values = set(file_loaders.values()) if file_loaders else set() diff --git a/pyproject.toml b/pyproject.toml index 117f09f8c..aeef6b67e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -9,7 +9,7 @@ dependencies = [ "chainlit>=2.2.1", "docling>=2.24.0", "einops>=0.8.1", - "hydra-core>=1.3.2", + "pyyaml>=6.0", "langchain-community>=0.3.18", "langchain-core>=0.3.39", "langchain-experimental>=0.3.4", diff --git a/pytest.ini b/pytest.ini index da258b2a7..517ec6d42 100644 --- a/pytest.ini +++ b/pytest.ini @@ -6,7 +6,6 @@ python_files = test_*.py env = - CONFIG_PATH=./.hydra_config PROMPTS_DIR=./prompts/example1 LOG_DIR=./logs diff --git a/quick_start/docker-compose.yaml b/quick_start/docker-compose.yaml index 55088b8fe..6374372a1 100644 --- a/quick_start/docker-compose.yaml +++ b/quick_start/docker-compose.yaml @@ -9,7 +9,6 @@ x-openrag: &openrag_template context: . dockerfile: Dockerfile volumes: - # - ${CONFIG_VOLUME:-./.hydra_config}:/app/.hydra_config - ${DATA_VOLUME:-./data}:/app/data - ${MODEL_WEIGHTS_VOLUME:-~/.cache/huggingface}:/app/model_weights # Model weights for RAG # - ./openrag:/app/openrag # For dev mode diff --git a/uv.lock b/uv.lock index 94dce647c..7d38d511e 100644 --- a/uv.lock +++ b/uv.lock @@ -1,5 +1,5 @@ version = 1 -revision = 2 +revision = 3 requires-python = ">=3.12" resolution-markers = [ "python_full_version >= '3.13'", @@ -166,12 +166,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/50/6f/346beae0375df5f6907230bc63d557ef5d7659be49250ac5931a758322ae/anthropic-0.46.0-py3-none-any.whl", hash = "sha256:1445ec9be78d2de7ea51b4d5acd3574e414aea97ef903d0ecbb57bec806aaa49", size = 223228, upload-time = "2025-02-18T20:35:28.659Z" }, ] -[[package]] -name = "antlr4-python3-runtime" -version = "4.9.3" -source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/3e/38/7859ff46355f76f8d19459005ca000b6e7012f2f1ca597746cbcd1fbfe5e/antlr4-python3-runtime-4.9.3.tar.gz", hash = "sha256:f224469b4168294902bb1efa80a8bf7855f24c99aef99cbefc1bcd3cce77881b", size = 117034, upload-time = "2021-11-06T17:52:23.524Z" } - [[package]] name = "anyio" version = "4.9.0" @@ -1102,7 +1096,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/f3/94/ad0d435f7c48debe960c53b8f60fb41c2026b1d0fa4a99a1cb17c3461e09/greenlet-3.2.3-cp312-cp312-macosx_11_0_universal2.whl", hash = "sha256:25ad29caed5783d4bd7a85c9251c651696164622494c00802a139c00d639242d", size = 271992, upload-time = "2025-06-05T16:11:23.467Z" }, { url = "https://files.pythonhosted.org/packages/93/5d/7c27cf4d003d6e77749d299c7c8f5fd50b4f251647b5c2e97e1f20da0ab5/greenlet-3.2.3-cp312-cp312-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:88cd97bf37fe24a6710ec6a3a7799f3f81d9cd33317dcf565ff9950c83f55e0b", size = 638820, upload-time = "2025-06-05T16:38:52.882Z" }, { url = "https://files.pythonhosted.org/packages/c6/7e/807e1e9be07a125bb4c169144937910bf59b9d2f6d931578e57f0bce0ae2/greenlet-3.2.3-cp312-cp312-manylinux2014_ppc64le.manylinux_2_17_ppc64le.whl", hash = "sha256:baeedccca94880d2f5666b4fa16fc20ef50ba1ee353ee2d7092b383a243b0b0d", size = 653046, upload-time = "2025-06-05T16:41:36.343Z" }, - { url = "https://files.pythonhosted.org/packages/9d/ab/158c1a4ea1068bdbc78dba5a3de57e4c7aeb4e7fa034320ea94c688bfb61/greenlet-3.2.3-cp312-cp312-manylinux2014_s390x.manylinux_2_17_s390x.whl", hash = "sha256:be52af4b6292baecfa0f397f3edb3c6092ce071b499dd6fe292c9ac9f2c8f264", size = 647701, upload-time = "2025-06-05T16:48:19.604Z" }, { url = "https://files.pythonhosted.org/packages/cc/0d/93729068259b550d6a0288da4ff72b86ed05626eaf1eb7c0d3466a2571de/greenlet-3.2.3-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:0cc73378150b8b78b0c9fe2ce56e166695e67478550769536a6742dca3651688", size = 649747, upload-time = "2025-06-05T16:13:04.628Z" }, { url = "https://files.pythonhosted.org/packages/f6/f6/c82ac1851c60851302d8581680573245c8fc300253fc1ff741ae74a6c24d/greenlet-3.2.3-cp312-cp312-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:706d016a03e78df129f68c4c9b4c4f963f7d73534e48a24f5f5a7101ed13dbbb", size = 605461, upload-time = "2025-06-05T16:12:50.792Z" }, { url = "https://files.pythonhosted.org/packages/98/82/d022cf25ca39cf1200650fc58c52af32c90f80479c25d1cbf57980ec3065/greenlet-3.2.3-cp312-cp312-musllinux_1_1_aarch64.whl", hash = "sha256:419e60f80709510c343c57b4bb5a339d8767bf9aef9b8ce43f4f143240f88b7c", size = 1121190, upload-time = "2025-06-05T16:36:48.59Z" }, @@ -1111,7 +1104,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/b1/cf/f5c0b23309070ae93de75c90d29300751a5aacefc0a3ed1b1d8edb28f08b/greenlet-3.2.3-cp313-cp313-macosx_11_0_universal2.whl", hash = "sha256:500b8689aa9dd1ab26872a34084503aeddefcb438e2e7317b89b11eaea1901ad", size = 270732, upload-time = "2025-06-05T16:10:08.26Z" }, { url = "https://files.pythonhosted.org/packages/48/ae/91a957ba60482d3fecf9be49bc3948f341d706b52ddb9d83a70d42abd498/greenlet-3.2.3-cp313-cp313-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:a07d3472c2a93117af3b0136f246b2833fdc0b542d4a9799ae5f41c28323faef", size = 639033, upload-time = "2025-06-05T16:38:53.983Z" }, { url = "https://files.pythonhosted.org/packages/6f/df/20ffa66dd5a7a7beffa6451bdb7400d66251374ab40b99981478c69a67a8/greenlet-3.2.3-cp313-cp313-manylinux2014_ppc64le.manylinux_2_17_ppc64le.whl", hash = "sha256:8704b3768d2f51150626962f4b9a9e4a17d2e37c8a8d9867bbd9fa4eb938d3b3", size = 652999, upload-time = "2025-06-05T16:41:37.89Z" }, - { url = "https://files.pythonhosted.org/packages/51/b4/ebb2c8cb41e521f1d72bf0465f2f9a2fd803f674a88db228887e6847077e/greenlet-3.2.3-cp313-cp313-manylinux2014_s390x.manylinux_2_17_s390x.whl", hash = "sha256:5035d77a27b7c62db6cf41cf786cfe2242644a7a337a0e155c80960598baab95", size = 647368, upload-time = "2025-06-05T16:48:21.467Z" }, { url = "https://files.pythonhosted.org/packages/8e/6a/1e1b5aa10dced4ae876a322155705257748108b7fd2e4fae3f2a091fe81a/greenlet-3.2.3-cp313-cp313-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:2d8aa5423cd4a396792f6d4580f88bdc6efcb9205891c9d40d20f6e670992efb", size = 650037, upload-time = "2025-06-05T16:13:06.402Z" }, { url = "https://files.pythonhosted.org/packages/26/f2/ad51331a157c7015c675702e2d5230c243695c788f8f75feba1af32b3617/greenlet-3.2.3-cp313-cp313-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:2c724620a101f8170065d7dded3f962a2aea7a7dae133a009cada42847e04a7b", size = 608402, upload-time = "2025-06-05T16:12:51.91Z" }, { url = "https://files.pythonhosted.org/packages/26/bc/862bd2083e6b3aff23300900a956f4ea9a4059de337f5c8734346b9b34fc/greenlet-3.2.3-cp313-cp313-musllinux_1_1_aarch64.whl", hash = "sha256:873abe55f134c48e1f2a6f53f7d1419192a3d1a4e873bace00499a4e45ea6af0", size = 1119577, upload-time = "2025-06-05T16:36:49.787Z" }, @@ -1120,7 +1112,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/d8/ca/accd7aa5280eb92b70ed9e8f7fd79dc50a2c21d8c73b9a0856f5b564e222/greenlet-3.2.3-cp314-cp314-macosx_11_0_universal2.whl", hash = "sha256:3d04332dddb10b4a211b68111dabaee2e1a073663d117dc10247b5b1642bac86", size = 271479, upload-time = "2025-06-05T16:10:47.525Z" }, { url = "https://files.pythonhosted.org/packages/55/71/01ed9895d9eb49223280ecc98a557585edfa56b3d0e965b9fa9f7f06b6d9/greenlet-3.2.3-cp314-cp314-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:8186162dffde068a465deab08fc72c767196895c39db26ab1c17c0b77a6d8b97", size = 683952, upload-time = "2025-06-05T16:38:55.125Z" }, { url = "https://files.pythonhosted.org/packages/ea/61/638c4bdf460c3c678a0a1ef4c200f347dff80719597e53b5edb2fb27ab54/greenlet-3.2.3-cp314-cp314-manylinux2014_ppc64le.manylinux_2_17_ppc64le.whl", hash = "sha256:f4bfbaa6096b1b7a200024784217defedf46a07c2eee1a498e94a1b5f8ec5728", size = 696917, upload-time = "2025-06-05T16:41:38.959Z" }, - { url = "https://files.pythonhosted.org/packages/22/cc/0bd1a7eb759d1f3e3cc2d1bc0f0b487ad3cc9f34d74da4b80f226fde4ec3/greenlet-3.2.3-cp314-cp314-manylinux2014_s390x.manylinux_2_17_s390x.whl", hash = "sha256:ed6cfa9200484d234d8394c70f5492f144b20d4533f69262d530a1a082f6ee9a", size = 692443, upload-time = "2025-06-05T16:48:23.113Z" }, { url = "https://files.pythonhosted.org/packages/67/10/b2a4b63d3f08362662e89c103f7fe28894a51ae0bc890fabf37d1d780e52/greenlet-3.2.3-cp314-cp314-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:02b0df6f63cd15012bed5401b47829cfd2e97052dc89da3cfaf2c779124eb892", size = 692995, upload-time = "2025-06-05T16:13:07.972Z" }, { url = "https://files.pythonhosted.org/packages/5a/c6/ad82f148a4e3ce9564056453a71529732baf5448ad53fc323e37efe34f66/greenlet-3.2.3-cp314-cp314-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:86c2d68e87107c1792e2e8d5399acec2487a4e993ab76c792408e59394d52141", size = 655320, upload-time = "2025-06-05T16:12:53.453Z" }, { url = "https://files.pythonhosted.org/packages/5c/4f/aab73ecaa6b3086a4c89863d94cf26fa84cbff63f52ce9bc4342b3087a06/greenlet-3.2.3-cp314-cp314-win_amd64.whl", hash = "sha256:8c47aae8fbbfcf82cc13327ae802ba13c9c36753b67e760023fd116bc124a62a", size = 301236, upload-time = "2025-06-05T16:15:20.111Z" }, @@ -1301,20 +1292,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/f0/0f/310fb31e39e2d734ccaa2c0fb981ee41f7bd5056ce9bc29b2248bd569169/humanfriendly-10.0-py2.py3-none-any.whl", hash = "sha256:1697e1a8a8f550fd43c2865cd84542fc175a61dcb779b6fee18cf6b6ccba1477", size = 86794, upload-time = "2021-09-17T21:40:39.897Z" }, ] -[[package]] -name = "hydra-core" -version = "1.3.2" -source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "antlr4-python3-runtime" }, - { name = "omegaconf" }, - { name = "packaging" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/6d/8e/07e42bc434a847154083b315779b0a81d567154504624e181caf2c71cd98/hydra-core-1.3.2.tar.gz", hash = "sha256:8a878ed67216997c3e9d88a8e72e7b4767e81af37afb4ea3334b269a4390a824", size = 3263494, upload-time = "2023-02-23T18:33:43.03Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/c6/50/e0edd38dcd63fb26a8547f13d28f7a008bc4a3fd4eb4ff030673f22ad41a/hydra_core-1.3.2-py3-none-any.whl", hash = "sha256:fa0238a9e31df3373b35b0bfb672c34cc92718d21f81311d8996a16de1141d8b", size = 154547, upload-time = "2023-02-23T18:33:40.801Z" }, -] - [[package]] name = "hyperframe" version = "6.1.0" @@ -2426,19 +2403,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/9e/4e/0d0c945463719429b7bd21dece907ad0bde437a2ff12b9b12fee94722ab0/nvidia_nvtx_cu12-12.6.77-py3-none-manylinux2014_x86_64.whl", hash = "sha256:6574241a3ec5fdc9334353ab8c479fe75841dbe8f4532a8fc97ce63503330ba1", size = 89265, upload-time = "2024-10-01T17:00:38.172Z" }, ] -[[package]] -name = "omegaconf" -version = "2.3.0" -source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "antlr4-python3-runtime" }, - { name = "pyyaml" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/09/48/6388f1bb9da707110532cb70ec4d2822858ddfb44f1cdf1233c20a80ea4b/omegaconf-2.3.0.tar.gz", hash = "sha256:d5d4b6d29955cc50ad50c46dc269bcd92c6e00f5f90d23ab5fee7bfca4ba4cc7", size = 3298120, upload-time = "2022-12-08T20:59:22.753Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/e3/94/1843518e420fa3ed6919835845df698c7e27e183cb997394e4a670973a65/omegaconf-2.3.0-py3-none-any.whl", hash = "sha256:7b4df175cdb08ba400f45cae3bdcae7ba8365db4d165fc65fd04b050ab63b46b", size = 79500, upload-time = "2022-12-08T20:59:19.686Z" }, -] - [[package]] name = "onnxruntime" version = "1.20.1" @@ -2555,7 +2519,6 @@ dependencies = [ { name = "faster-whisper" }, { name = "hdbscan" }, { name = "html-to-markdown" }, - { name = "hydra-core" }, { name = "infinity-client" }, { name = "langchain-community" }, { name = "langchain-core" }, @@ -2579,6 +2542,7 @@ dependencies = [ { name = "pymupdf4llm" }, { name = "pytest-env" }, { name = "python-dotenv" }, + { name = "pyyaml" }, { name = "ray", extra = ["default"] }, { name = "robotframework" }, { name = "robotframework-requests" }, @@ -2614,7 +2578,6 @@ requires-dist = [ { name = "faster-whisper", specifier = ">=1.1.0" }, { name = "hdbscan", specifier = ">=0.8.40" }, { name = "html-to-markdown", specifier = ">=2.4.0" }, - { name = "hydra-core", specifier = ">=1.3.2" }, { name = "infinity-client", specifier = ">=0.0.76" }, { name = "langchain-community", specifier = ">=0.3.18" }, { name = "langchain-core", specifier = ">=0.3.39" }, @@ -2638,6 +2601,7 @@ requires-dist = [ { name = "pymupdf4llm", specifier = ">=0.0.17" }, { name = "pytest-env", specifier = ">=1.1.5" }, { name = "python-dotenv", specifier = ">=1.0.1" }, + { name = "pyyaml", specifier = ">=6.0" }, { name = "ray", extras = ["default"], specifier = ">=2.47.1" }, { name = "robotframework", specifier = ">=7.2.2" }, { name = "robotframework-requests", specifier = ">=0.9.7" },