Drop ConfigProxy in favour of passing AppConfig as an keyword param to HaikuRag clients.

This commit is contained in:
Yiorgis Gozadinos 2025-10-27 10:33:33 +02:00
parent 110cc447e7
commit 9c8ffcc212
No known key found for this signature in database
9 changed files with 114 additions and 111 deletions

View file

@ -104,10 +104,10 @@ a2a:
## Programmatic Configuration ## Programmatic Configuration
When using haiku.rag as a Python library, you can configure it programmatically using `set_config()`: When using haiku.rag as a Python library, you can pass configuration directly to the `HaikuRAG` client:
```python ```python
from haiku.rag.config import set_config, AppConfig from haiku.rag.config import AppConfig
from haiku.rag.client import HaikuRAG from haiku.rag.client import HaikuRAG
# Create custom configuration # Create custom configuration
@ -117,18 +117,17 @@ custom_config = AppConfig(
processing={"chunk_size": 512} processing={"chunk_size": 512}
) )
# Set the configuration globally # Pass configuration to the client
set_config(custom_config) client = HaikuRAG(config=custom_config)
# All subsequent operations use this configuration
client = HaikuRAG(db_path)
``` ```
If you don't pass a config, the client uses the global configuration loaded from your YAML file or defaults.
This is useful for: This is useful for:
- Jupyter notebooks - Jupyter notebooks
- Python scripts - Python scripts
- Testing with different configurations - Testing with different configurations
- Applications that need runtime configuration - Applications that need multiple clients with different configurations
## API Keys ## API Keys

View file

@ -8,7 +8,7 @@ from urllib.parse import urlparse
import httpx import httpx
from haiku.rag.config import Config from haiku.rag.config import AppConfig, Config
from haiku.rag.reranking import get_reranker from haiku.rag.reranking import get_reranker
from haiku.rag.store.engine import Store from haiku.rag.store.engine import Store
from haiku.rag.store.models.chunk import Chunk from haiku.rag.store.models.chunk import Chunk
@ -25,16 +25,23 @@ class HaikuRAG:
def __init__( def __init__(
self, self,
db_path: Path = Config.storage.data_dir / "haiku.rag.lancedb", db_path: Path | None = None,
config: AppConfig = Config,
skip_validation: bool = False, skip_validation: bool = False,
): ):
"""Initialize the RAG client with a database path. """Initialize the RAG client with a database path.
Args: Args:
db_path: Path to the database file. db_path: Path to the database file. If None, uses config.storage.data_dir.
config: Configuration to use. Defaults to global Config.
skip_validation: Whether to skip configuration validation on database load. skip_validation: Whether to skip configuration validation on database load.
""" """
self.store = Store(db_path, skip_validation=skip_validation) self._config = config
if db_path is None:
db_path = self._config.storage.data_dir / "haiku.rag.lancedb"
self.store = Store(
db_path, config=self._config, skip_validation=skip_validation
)
self.document_repository = DocumentRepository(self.store) self.document_repository = DocumentRepository(self.store)
self.chunk_repository = ChunkRepository(self.store) self.chunk_repository = ChunkRepository(self.store)
@ -430,7 +437,7 @@ class HaikuRAG:
List of (chunk, score) tuples ordered by relevance. List of (chunk, score) tuples ordered by relevance.
""" """
# Get reranker if available # Get reranker if available
reranker = get_reranker() reranker = get_reranker(config=self._config)
if reranker is None: if reranker is None:
# No reranking - return direct search results # No reranking - return direct search results
@ -452,18 +459,20 @@ class HaikuRAG:
async def expand_context( async def expand_context(
self, self,
search_results: list[tuple[Chunk, float]], search_results: list[tuple[Chunk, float]],
radius: int = Config.processing.context_chunk_radius, radius: int | None = None,
) -> list[tuple[Chunk, float]]: ) -> list[tuple[Chunk, float]]:
"""Expand search results with adjacent chunks, merging overlapping chunks. """Expand search results with adjacent chunks, merging overlapping chunks.
Args: Args:
search_results: List of (chunk, score) tuples from search. search_results: List of (chunk, score) tuples from search.
radius: Number of adjacent chunks to include before/after each chunk. radius: Number of adjacent chunks to include before/after each chunk.
Defaults to CONTEXT_CHUNK_RADIUS config setting. If None, uses config.processing.context_chunk_radius.
Returns: Returns:
List of (chunk, score) tuples with expanded and merged context chunks. List of (chunk, score) tuples with expanded and merged context chunks.
""" """
if radius is None:
radius = self._config.processing.context_chunk_radius
if radius == 0: if radius == 0:
return search_results return search_results
@ -593,7 +602,9 @@ class HaikuRAG:
""" """
from haiku.rag.qa import get_qa_agent from haiku.rag.qa import get_qa_agent
qa_agent = get_qa_agent(self, use_citations=cite, system_prompt=system_prompt) qa_agent = get_qa_agent(
self, config=self._config, use_citations=cite, system_prompt=system_prompt
)
return await qa_agent.answer(question) return await qa_agent.answer(question)
async def rebuild_database(self) -> AsyncGenerator[str, None]: async def rebuild_database(self) -> AsyncGenerator[str, None]:

View file

@ -40,53 +40,15 @@ __all__ = [
"load_yaml_config", "load_yaml_config",
"generate_default_config", "generate_default_config",
"load_config_from_env", "load_config_from_env",
"set_config",
] ]
# Load config from YAML file or use defaults
class ConfigProxy: config_path = find_config_file(None)
"""Proxy for the global configuration that allows runtime updates.""" if config_path:
yaml_data = load_yaml_config(config_path)
def __init__(self): Config = AppConfig.model_validate(yaml_data)
# Load config from YAML file or use defaults else:
config_path = find_config_file(None) Config = AppConfig()
if config_path:
yaml_data = load_yaml_config(config_path)
self._config = AppConfig.model_validate(yaml_data)
else:
self._config = AppConfig()
def __getattr__(self, name):
"""Proxy attribute access to the underlying config."""
return getattr(self._config, name)
def set(self, config: AppConfig) -> None:
"""Replace the current configuration."""
self._config = config
# Create the global Config instance
Config = ConfigProxy()
# Check for deprecated .env file # Check for deprecated .env file
check_for_deprecated_env() check_for_deprecated_env()
def set_config(config: AppConfig) -> None:
"""Set the global configuration programmatically.
This allows library users to configure haiku.rag without needing
a YAML file or environment variables.
Args:
config: The AppConfig instance to use globally.
Example:
>>> from haiku.rag.config import set_config, AppConfig
>>> custom_config = AppConfig(
... qa={"provider": "openai", "model": "gpt-4o"},
... embeddings={"provider": "voyage", "model": "voyage-3"}
... )
>>> set_config(custom_config)
"""
Config.set(config)

View file

@ -1,17 +1,23 @@
from haiku.rag.config import Config from haiku.rag.config import AppConfig, Config
from haiku.rag.embeddings.base import EmbedderBase from haiku.rag.embeddings.base import EmbedderBase
from haiku.rag.embeddings.ollama import Embedder as OllamaEmbedder from haiku.rag.embeddings.ollama import Embedder as OllamaEmbedder
def get_embedder() -> EmbedderBase: def get_embedder(config: AppConfig = Config) -> EmbedderBase:
""" """
Factory function to get the appropriate embedder based on the configuration. Factory function to get the appropriate embedder based on the configuration.
Args:
config: Configuration to use. Defaults to global Config.
Returns:
An embedder instance configured according to the config.
""" """
if Config.embeddings.provider == "ollama": if config.embeddings.provider == "ollama":
return OllamaEmbedder(Config.embeddings.model, Config.embeddings.vector_dim) return OllamaEmbedder(config.embeddings.model, config.embeddings.vector_dim)
if Config.embeddings.provider == "voyageai": if config.embeddings.provider == "voyageai":
try: try:
from haiku.rag.embeddings.voyageai import Embedder as VoyageAIEmbedder from haiku.rag.embeddings.voyageai import Embedder as VoyageAIEmbedder
except ImportError: except ImportError:
@ -20,16 +26,16 @@ def get_embedder() -> EmbedderBase:
"Please install haiku.rag with the 'voyageai' extra: " "Please install haiku.rag with the 'voyageai' extra: "
"uv pip install haiku.rag[voyageai]" "uv pip install haiku.rag[voyageai]"
) )
return VoyageAIEmbedder(Config.embeddings.model, Config.embeddings.vector_dim) return VoyageAIEmbedder(config.embeddings.model, config.embeddings.vector_dim)
if Config.embeddings.provider == "openai": if config.embeddings.provider == "openai":
from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder
return OpenAIEmbedder(Config.embeddings.model, Config.embeddings.vector_dim) return OpenAIEmbedder(config.embeddings.model, config.embeddings.vector_dim)
if Config.embeddings.provider == "vllm": if config.embeddings.provider == "vllm":
from haiku.rag.embeddings.vllm import Embedder as VllmEmbedder from haiku.rag.embeddings.vllm import Embedder as VllmEmbedder
return VllmEmbedder(Config.embeddings.model, Config.embeddings.vector_dim) return VllmEmbedder(config.embeddings.model, config.embeddings.vector_dim)
raise ValueError(f"Unsupported embedding provider: {Config.embeddings.provider}") raise ValueError(f"Unsupported embedding provider: {config.embeddings.provider}")

View file

@ -1,15 +1,28 @@
from haiku.rag.client import HaikuRAG from haiku.rag.client import HaikuRAG
from haiku.rag.config import Config from haiku.rag.config import AppConfig, Config
from haiku.rag.qa.agent import QuestionAnswerAgent from haiku.rag.qa.agent import QuestionAnswerAgent
def get_qa_agent( def get_qa_agent(
client: HaikuRAG, client: HaikuRAG,
config: AppConfig = Config,
use_citations: bool = False, use_citations: bool = False,
system_prompt: str | None = None, system_prompt: str | None = None,
) -> QuestionAnswerAgent: ) -> QuestionAnswerAgent:
provider = Config.qa.provider """
model_name = Config.qa.model Factory function to get a QA agent based on the configuration.
Args:
client: HaikuRAG client instance.
config: Configuration to use. Defaults to global Config.
use_citations: Whether to include citations in responses.
system_prompt: Optional custom system prompt.
Returns:
A configured QuestionAnswerAgent instance.
"""
provider = config.qa.provider
model_name = config.qa.model
return QuestionAnswerAgent( return QuestionAnswerAgent(
client=client, client=client,

View file

@ -1,37 +1,45 @@
import os import os
from haiku.rag.config import Config from haiku.rag.config import AppConfig, Config
from haiku.rag.reranking.base import RerankerBase from haiku.rag.reranking.base import RerankerBase
_reranker: RerankerBase | None = None _reranker_cache: dict[int, RerankerBase | None] = {}
def get_reranker() -> RerankerBase | None: def get_reranker(config: AppConfig = Config) -> RerankerBase | None:
""" """
Factory function to get the appropriate reranker based on the configuration. Factory function to get the appropriate reranker based on the configuration.
Returns None if if reranking is disabled. Returns None if reranking is disabled.
"""
global _reranker
if _reranker is not None:
return _reranker
if Config.reranking.provider == "mxbai": Args:
config: Configuration to use. Defaults to global Config.
Returns:
A reranker instance if configured, None otherwise.
"""
# Use config id as cache key to support multiple configs
config_id = id(config)
if config_id in _reranker_cache:
return _reranker_cache[config_id]
reranker: RerankerBase | None = None
if config.reranking.provider == "mxbai":
try: try:
from haiku.rag.reranking.mxbai import MxBAIReranker from haiku.rag.reranking.mxbai import MxBAIReranker
os.environ["TOKENIZERS_PARALLELISM"] = "true" os.environ["TOKENIZERS_PARALLELISM"] = "true"
_reranker = MxBAIReranker() reranker = MxBAIReranker()
return _reranker
except ImportError: except ImportError:
return None reranker = None
if Config.reranking.provider == "cohere": elif config.reranking.provider == "cohere":
try: try:
from haiku.rag.reranking.cohere import CohereReranker from haiku.rag.reranking.cohere import CohereReranker
_reranker = CohereReranker() reranker = CohereReranker()
return _reranker
except ImportError: except ImportError:
return None reranker = None
return None _reranker_cache[config_id] = reranker
return reranker

View file

@ -10,7 +10,7 @@ import lancedb
from lancedb.pydantic import LanceModel, Vector from lancedb.pydantic import LanceModel, Vector
from pydantic import Field from pydantic import Field
from haiku.rag.config import Config from haiku.rag.config import AppConfig, Config
from haiku.rag.embeddings import get_embedder from haiku.rag.embeddings import get_embedder
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@ -49,9 +49,12 @@ class SettingsRecord(LanceModel):
class Store: class Store:
def __init__(self, db_path: Path, skip_validation: bool = False): def __init__(
self, db_path: Path, config: AppConfig = Config, skip_validation: bool = False
):
self.db_path: Path = db_path self.db_path: Path = db_path
self.embedder = get_embedder() self._config = config
self.embedder = get_embedder(config=self._config)
self._vacuum_lock = asyncio.Lock() self._vacuum_lock = asyncio.Lock()
# Create the ChunkRecord model with the correct vector dimension # Create the ChunkRecord model with the correct vector dimension
@ -59,7 +62,7 @@ class Store:
# Local filesystem handling for DB directory # Local filesystem handling for DB directory
if not self._has_cloud_config(): if not self._has_cloud_config():
if Config.storage.disable_autocreate: if self._config.storage.disable_autocreate:
# LanceDB uses a directory path for local databases; enforce presence # LanceDB uses a directory path for local databases; enforce presence
if not db_path.exists(): if not db_path.exists():
raise FileNotFoundError( raise FileNotFoundError(
@ -85,13 +88,15 @@ class Store:
Args: Args:
retention_seconds: Retention threshold in seconds. Only versions older retention_seconds: Retention threshold in seconds. Only versions older
than this will be removed. If None, uses Config.storage.vacuum_retention_seconds. than this will be removed. If None, uses config.storage.vacuum_retention_seconds.
Note: Note:
If vacuum is already running, this method returns immediately without blocking. If vacuum is already running, this method returns immediately without blocking.
Use asyncio.create_task(store.vacuum()) for non-blocking background execution. Use asyncio.create_task(store.vacuum()) for non-blocking background execution.
""" """
if self._has_cloud_config() and str(Config.lancedb.uri).startswith("db://"): if self._has_cloud_config() and str(self._config.lancedb.uri).startswith(
"db://"
):
return return
# Skip if already running (non-blocking) # Skip if already running (non-blocking)
@ -102,7 +107,7 @@ class Store:
try: try:
# Evaluate config at runtime to allow dynamic changes # Evaluate config at runtime to allow dynamic changes
if retention_seconds is None: if retention_seconds is None:
retention_seconds = Config.storage.vacuum_retention_seconds retention_seconds = self._config.storage.vacuum_retention_seconds
# Perform maintenance per table using optimize() with configurable retention # Perform maintenance per table using optimize() with configurable retention
retention = timedelta(seconds=retention_seconds) retention = timedelta(seconds=retention_seconds)
for table in [ for table in [
@ -120,9 +125,9 @@ class Store:
# Check if we have cloud configuration # Check if we have cloud configuration
if self._has_cloud_config(): if self._has_cloud_config():
return lancedb.connect( return lancedb.connect(
uri=Config.lancedb.uri, uri=self._config.lancedb.uri,
api_key=Config.lancedb.api_key, api_key=self._config.lancedb.api_key,
region=Config.lancedb.region, region=self._config.lancedb.region,
) )
else: else:
# Local file system connection # Local file system connection
@ -131,7 +136,9 @@ class Store:
def _has_cloud_config(self) -> bool: def _has_cloud_config(self) -> bool:
"""Check if cloud configuration is complete.""" """Check if cloud configuration is complete."""
return bool( return bool(
Config.lancedb.uri and Config.lancedb.api_key and Config.lancedb.region self._config.lancedb.uri
and self._config.lancedb.api_key
and self._config.lancedb.region
) )
def _validate_configuration(self) -> None: def _validate_configuration(self) -> None:
@ -173,7 +180,7 @@ class Store:
"settings", schema=SettingsRecord "settings", schema=SettingsRecord
) )
# Save current settings to the new database # Save current settings to the new database
settings_data = Config.model_dump(mode="json") settings_data = self._config.model_dump(mode="json")
self.settings_table.add( self.settings_table.add(
[SettingsRecord(id="settings", settings=json.dumps(settings_data))] [SettingsRecord(id="settings", settings=json.dumps(settings_data))]
) )

View file

@ -6,8 +6,6 @@ from uuid import uuid4
from lancedb.rerankers import RRFReranker from lancedb.rerankers import RRFReranker
from haiku.rag.config import Config
from haiku.rag.embeddings import get_embedder
from haiku.rag.store.engine import DocumentRecord, Store from haiku.rag.store.engine import DocumentRecord, Store
from haiku.rag.store.models.chunk import Chunk from haiku.rag.store.models.chunk import Chunk
from haiku.rag.utils import load_callable from haiku.rag.utils import load_callable
@ -23,7 +21,7 @@ class ChunkRepository:
def __init__(self, store: Store) -> None: def __init__(self, store: Store) -> None:
self.store = store self.store = store
self.embedder = get_embedder() self.embedder = store.embedder
def _ensure_fts_index(self) -> None: def _ensure_fts_index(self) -> None:
"""Ensure FTS index exists on the content column.""" """Ensure FTS index exists on the content column."""
@ -153,7 +151,7 @@ class ChunkRepository:
# Optionally preprocess markdown before chunking # Optionally preprocess markdown before chunking
processed_document = document processed_document = document
preprocessor_path = Config.processing.markdown_preprocessor preprocessor_path = self.store._config.processing.markdown_preprocessor
if preprocessor_path: if preprocessor_path:
try: try:
pre_fn = load_callable(preprocessor_path) pre_fn = load_callable(preprocessor_path)

View file

@ -1,6 +1,5 @@
import json import json
from haiku.rag.config import Config
from haiku.rag.store.engine import SettingsRecord, Store from haiku.rag.store.engine import SettingsRecord, Store
@ -73,7 +72,7 @@ class SettingsRepository:
def save_current_settings(self) -> None: def save_current_settings(self) -> None:
"""Save the current configuration to the database.""" """Save the current configuration to the database."""
current_config = Config.model_dump(mode="json") current_config = self.store._config.model_dump(mode="json")
# Check if settings exist # Check if settings exist
existing = list( existing = list(
@ -116,7 +115,7 @@ class SettingsRepository:
self.save_current_settings() self.save_current_settings()
return return
current_config = Config.model_dump(mode="json") current_config = self.store._config.model_dump(mode="json")
# Check if embedding provider or model has changed # Check if embedding provider or model has changed
# Support both old flat structure and new nested structure for backward compatibility # Support both old flat structure and new nested structure for backward compatibility