Merge pull request #114 from ggozad/feat/yaml-config
Configure haiku.rag with YAML instead of env vars.
This commit is contained in:
commit
afd4a581fb
62 changed files with 1339 additions and 497 deletions
|
|
@ -25,6 +25,13 @@ wheels/
|
||||||
venv/
|
venv/
|
||||||
env/
|
env/
|
||||||
|
|
||||||
|
# Node.js
|
||||||
|
node_modules/
|
||||||
|
.next/
|
||||||
|
npm-debug.log*
|
||||||
|
yarn-debug.log*
|
||||||
|
yarn-error.log*
|
||||||
|
|
||||||
# Data
|
# Data
|
||||||
*.lancedb/
|
*.lancedb/
|
||||||
data/
|
data/
|
||||||
|
|
|
||||||
5
.gitignore
vendored
5
.gitignore
vendored
|
|
@ -16,8 +16,9 @@ tests/data/
|
||||||
.pytest_cache/
|
.pytest_cache/
|
||||||
.ruff_cache/
|
.ruff_cache/
|
||||||
|
|
||||||
# environment variables
|
# environment variables and config files
|
||||||
.env
|
.env
|
||||||
|
haiku.rag.yaml
|
||||||
TODO.md
|
TODO.md
|
||||||
PLAN.md
|
PLAN.md
|
||||||
DEVNOTES.md
|
DEVNOTES.md
|
||||||
|
|
@ -27,4 +28,4 @@ DEVNOTES.md
|
||||||
.mcpregistry_registry_token
|
.mcpregistry_registry_token
|
||||||
|
|
||||||
# MkDocs site directory when doing local docs builds
|
# MkDocs site directory when doing local docs builds
|
||||||
site/
|
site/
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,8 @@ Retrieval-Augmented Generation (RAG) library built on LanceDB.
|
||||||
|
|
||||||
`haiku.rag` is a Retrieval-Augmented Generation (RAG) library built to work with LanceDB as a local vector database. It uses LanceDB for storing embeddings and performs semantic (vector) search as well as full-text search combined through native hybrid search with Reciprocal Rank Fusion. Both open-source (Ollama) as well as commercial (OpenAI, VoyageAI) embedding providers are supported.
|
`haiku.rag` is a Retrieval-Augmented Generation (RAG) library built to work with LanceDB as a local vector database. It uses LanceDB for storing embeddings and performs semantic (vector) search as well as full-text search combined through native hybrid search with Reciprocal Rank Fusion. Both open-source (Ollama) as well as commercial (OpenAI, VoyageAI) embedding providers are supported.
|
||||||
|
|
||||||
|
> **Note**: Configuration now uses YAML files instead of environment variables. If you're upgrading from an older version, run `haiku-rag init-config --from-env` to migrate your `.env` file to `haiku.rag.yaml`. See [Configuration](https://ggozad.github.io/haiku.rag/configuration/) for details.
|
||||||
|
|
||||||
## Features
|
## Features
|
||||||
|
|
||||||
- **Local LanceDB**: No external servers required, supports also LanceDB cloud storage, S3, Google Cloud & Azure
|
- **Local LanceDB**: No external servers required, supports also LanceDB cloud storage, S3, Google Cloud & Azure
|
||||||
|
|
@ -58,10 +60,11 @@ haiku-rag research \
|
||||||
haiku-rag rebuild
|
haiku-rag rebuild
|
||||||
|
|
||||||
# Start server with file monitoring
|
# Start server with file monitoring
|
||||||
export MONITOR_DIRECTORIES="/path/to/docs"
|
haiku-rag serve --monitor
|
||||||
haiku-rag serve
|
|
||||||
```
|
```
|
||||||
|
|
||||||
|
To customize settings, create a `haiku.rag.yaml` config file (see [Configuration](https://ggozad.github.io/haiku.rag/configuration/)).
|
||||||
|
|
||||||
## Python Usage
|
## Python Usage
|
||||||
|
|
||||||
```python
|
```python
|
||||||
|
|
@ -172,7 +175,7 @@ See the [examples directory](examples/) for working examples:
|
||||||
Full documentation at: https://ggozad.github.io/haiku.rag/
|
Full documentation at: https://ggozad.github.io/haiku.rag/
|
||||||
|
|
||||||
- [Installation](https://ggozad.github.io/haiku.rag/installation/) - Provider setup
|
- [Installation](https://ggozad.github.io/haiku.rag/installation/) - Provider setup
|
||||||
- [Configuration](https://ggozad.github.io/haiku.rag/configuration/) - Environment variables
|
- [Configuration](https://ggozad.github.io/haiku.rag/configuration/) - YAML configuration
|
||||||
- [CLI](https://ggozad.github.io/haiku.rag/cli/) - Command reference
|
- [CLI](https://ggozad.github.io/haiku.rag/cli/) - Command reference
|
||||||
- [Python API](https://ggozad.github.io/haiku.rag/python/) - Complete API docs
|
- [Python API](https://ggozad.github.io/haiku.rag/python/) - Complete API docs
|
||||||
- [Agents](https://ggozad.github.io/haiku.rag/agents/) - QA agent and multi-agent research
|
- [Agents](https://ggozad.github.io/haiku.rag/agents/) - QA agent and multi-agent research
|
||||||
|
|
|
||||||
|
|
@ -8,15 +8,46 @@ Pre-built images are available at `ghcr.io/ggozad/haiku.rag` with all extras (vo
|
||||||
docker pull ghcr.io/ggozad/haiku.rag:latest
|
docker pull ghcr.io/ggozad/haiku.rag:latest
|
||||||
```
|
```
|
||||||
|
|
||||||
|
## Configuration
|
||||||
|
|
||||||
|
Create a configuration file `haiku.rag.yaml`:
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
# haiku.rag.yaml
|
||||||
|
environment: production
|
||||||
|
|
||||||
|
embeddings:
|
||||||
|
provider: ollama
|
||||||
|
model: nomic-embed-text
|
||||||
|
vector_dim: 768
|
||||||
|
|
||||||
|
qa:
|
||||||
|
provider: ollama
|
||||||
|
model: qwen3
|
||||||
|
```
|
||||||
|
|
||||||
|
See [Configuration docs](https://ggozad.github.io/haiku.rag/configuration/) for all available options.
|
||||||
|
|
||||||
## Running
|
## Running
|
||||||
|
|
||||||
|
Mount your config file and data directory:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
docker run -p 8000:8000 -p 8001:8001 \
|
docker run -p 8000:8000 -p 8001:8001 \
|
||||||
|
-v $(pwd)/haiku.rag.yaml:/app/haiku.rag.yaml \
|
||||||
-v $(pwd)/data:/data \
|
-v $(pwd)/data:/data \
|
||||||
-e EMBEDDINGS_PROVIDER=ollama \
|
ghcr.io/ggozad/haiku.rag:latest
|
||||||
-e EMBEDDINGS_MODEL=nomic-embed-text \
|
```
|
||||||
-e QA_PROVIDER=ollama \
|
|
||||||
-e QA_MODEL=qwen3\
|
The container will automatically use the mounted `haiku.rag.yaml` configuration file.
|
||||||
|
|
||||||
|
For API keys (OpenAI, Anthropic, etc.), pass them as environment variables:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
docker run -p 8000:8000 -p 8001:8001 \
|
||||||
|
-v $(pwd)/haiku.rag.yaml:/app/haiku.rag.yaml \
|
||||||
|
-v $(pwd)/data:/data \
|
||||||
|
-e OPENAI_API_KEY=your-key-here \
|
||||||
ghcr.io/ggozad/haiku.rag:latest
|
ghcr.io/ggozad/haiku.rag:latest
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -26,8 +57,6 @@ docker run -p 8000:8000 -p 8001:8001 \
|
||||||
docker build -f docker/Dockerfile -t haiku-rag .
|
docker build -f docker/Dockerfile -t haiku-rag .
|
||||||
```
|
```
|
||||||
|
|
||||||
Note: The environment variables above override the defaults. See [Configuration docs](https://ggozad.github.io/haiku.rag/configuration/) for all options.
|
|
||||||
|
|
||||||
## Docker Compose
|
## Docker Compose
|
||||||
|
|
||||||
See `examples/docker/` for a complete setup example.
|
See `examples/docker/` for a complete setup example.
|
||||||
|
|
|
||||||
|
|
@ -144,12 +144,13 @@ All operations create artifacts for traceability:
|
||||||
|
|
||||||
To prevent memory growth, the server uses LRU (Least Recently Used) eviction:
|
To prevent memory growth, the server uses LRU (Least Recently Used) eviction:
|
||||||
|
|
||||||
- Maximum 1000 contexts kept in memory (configurable via `A2A_MAX_CONTEXTS`)
|
- Maximum 1000 contexts kept in memory (configurable via `a2a.max_contexts`)
|
||||||
- When limit exceeded, least recently used contexts are automatically evicted
|
- When limit exceeded, least recently used contexts are automatically evicted
|
||||||
|
|
||||||
Configure via environment variable:
|
Configure in `haiku.rag.yaml`:
|
||||||
```bash
|
```yaml
|
||||||
export A2A_MAX_CONTEXTS=1000
|
a2a:
|
||||||
|
max_contexts: 1000
|
||||||
```
|
```
|
||||||
|
|
||||||
## Security
|
## Security
|
||||||
|
|
|
||||||
|
|
@ -1,20 +1,161 @@
|
||||||
# Configuration
|
# Configuration
|
||||||
|
|
||||||
Configuration is done through the use of environment variables.
|
Configuration is done through YAML configuration files.
|
||||||
|
|
||||||
!!! note
|
!!! note
|
||||||
If you create a db with certain settings and later change them, `haiku.rag` will detect incompatibilities (for example, if you change embedding provider) and will exit. You can **rebuild** the database to apply the new settings, see [Rebuild Database](./cli.md#rebuild-database).
|
If you create a db with certain settings and later change them, `haiku.rag` will detect incompatibilities (for example, if you change embedding provider) and will exit. You can **rebuild** the database to apply the new settings, see [Rebuild Database](./cli.md#rebuild-database).
|
||||||
|
|
||||||
|
## Getting Started
|
||||||
|
|
||||||
|
Generate a configuration file with defaults:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
haiku-rag init-config
|
||||||
|
```
|
||||||
|
|
||||||
|
This creates a `haiku.rag.yaml` file in your current directory with all available settings.
|
||||||
|
|
||||||
|
To migrate from environment variables (`.env` file):
|
||||||
|
|
||||||
|
```bash
|
||||||
|
haiku-rag init-config --from-env
|
||||||
|
```
|
||||||
|
|
||||||
|
## Configuration File Locations
|
||||||
|
|
||||||
|
`haiku.rag` searches for configuration files in this order:
|
||||||
|
|
||||||
|
1. Path specified via `--config` flag: `haiku-rag --config /path/to/config.yaml <command>`
|
||||||
|
2. `./haiku.rag.yaml` (current directory)
|
||||||
|
3. `~/.config/haiku.rag/config.yaml` (user config directory)
|
||||||
|
|
||||||
|
## Minimal Configuration
|
||||||
|
|
||||||
|
A minimal configuration file with defaults:
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
# haiku.rag.yaml
|
||||||
|
environment: production
|
||||||
|
|
||||||
|
embeddings:
|
||||||
|
provider: ollama
|
||||||
|
model: qwen3-embedding
|
||||||
|
vector_dim: 4096
|
||||||
|
|
||||||
|
qa:
|
||||||
|
provider: ollama
|
||||||
|
model: gpt-oss
|
||||||
|
```
|
||||||
|
|
||||||
|
## Complete Configuration Example
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
# haiku.rag.yaml
|
||||||
|
environment: production
|
||||||
|
|
||||||
|
storage:
|
||||||
|
data_dir: "" # Empty = use default platform location
|
||||||
|
monitor_directories:
|
||||||
|
- /path/to/documents
|
||||||
|
- /another/path
|
||||||
|
disable_autocreate: false
|
||||||
|
vacuum_retention_seconds: 60
|
||||||
|
|
||||||
|
lancedb:
|
||||||
|
uri: "" # Empty for local, or db://, s3://, az://, gs://
|
||||||
|
api_key: ""
|
||||||
|
region: ""
|
||||||
|
|
||||||
|
embeddings:
|
||||||
|
provider: ollama
|
||||||
|
model: qwen3-embedding
|
||||||
|
vector_dim: 4096
|
||||||
|
|
||||||
|
reranking:
|
||||||
|
provider: "" # Empty to disable, or mxbai, cohere, vllm
|
||||||
|
model: ""
|
||||||
|
|
||||||
|
qa:
|
||||||
|
provider: ollama
|
||||||
|
model: gpt-oss
|
||||||
|
|
||||||
|
research:
|
||||||
|
provider: "" # Empty to use qa settings
|
||||||
|
model: ""
|
||||||
|
|
||||||
|
processing:
|
||||||
|
chunk_size: 256
|
||||||
|
context_chunk_radius: 0
|
||||||
|
markdown_preprocessor: ""
|
||||||
|
|
||||||
|
providers:
|
||||||
|
ollama:
|
||||||
|
base_url: http://localhost:11434
|
||||||
|
|
||||||
|
vllm:
|
||||||
|
embeddings_base_url: ""
|
||||||
|
rerank_base_url: ""
|
||||||
|
qa_base_url: ""
|
||||||
|
research_base_url: ""
|
||||||
|
|
||||||
|
a2a:
|
||||||
|
max_contexts: 1000
|
||||||
|
```
|
||||||
|
|
||||||
|
## Programmatic Configuration
|
||||||
|
|
||||||
|
When using haiku.rag as a Python library, you can pass configuration directly to the `HaikuRAG` client:
|
||||||
|
|
||||||
|
```python
|
||||||
|
from haiku.rag.config import AppConfig
|
||||||
|
from haiku.rag.client import HaikuRAG
|
||||||
|
|
||||||
|
# Create custom configuration
|
||||||
|
custom_config = AppConfig(
|
||||||
|
qa={"provider": "openai", "model": "gpt-4o"},
|
||||||
|
embeddings={"provider": "ollama", "model": "qwen3-embedding"},
|
||||||
|
processing={"chunk_size": 512}
|
||||||
|
)
|
||||||
|
|
||||||
|
# Pass configuration to the client
|
||||||
|
client = HaikuRAG(config=custom_config)
|
||||||
|
```
|
||||||
|
|
||||||
|
If you don't pass a config, the client uses the global configuration loaded from your YAML file or defaults.
|
||||||
|
|
||||||
|
This is useful for:
|
||||||
|
- Jupyter notebooks
|
||||||
|
- Python scripts
|
||||||
|
- Testing with different configurations
|
||||||
|
- Applications that need multiple clients with different configurations
|
||||||
|
|
||||||
|
## API Keys
|
||||||
|
|
||||||
|
API keys are configured through **environment variables**, not in the YAML file.
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# OpenAI
|
||||||
|
export OPENAI_API_KEY=your-key-here
|
||||||
|
|
||||||
|
# Anthropic
|
||||||
|
export ANTHROPIC_API_KEY=your-key-here
|
||||||
|
|
||||||
|
# Voyage AI
|
||||||
|
export VOYAGE_API_KEY=your-key-here
|
||||||
|
|
||||||
|
# Cohere
|
||||||
|
export CO_API_KEY=your-key-here
|
||||||
|
```
|
||||||
|
|
||||||
## File Monitoring
|
## File Monitoring
|
||||||
|
|
||||||
Set directories to monitor for automatic indexing:
|
Set directories to monitor for automatic indexing:
|
||||||
|
|
||||||
```bash
|
```yaml
|
||||||
# Monitor single directory
|
storage:
|
||||||
MONITOR_DIRECTORIES="/path/to/documents"
|
monitor_directories:
|
||||||
|
- /path/to/documents
|
||||||
# Monitor multiple directories
|
- /another_path/to/documents
|
||||||
MONITOR_DIRECTORIES="/path/to/documents,/another_path/to/documents"
|
|
||||||
```
|
```
|
||||||
|
|
||||||
## Embedding Providers
|
## Embedding Providers
|
||||||
|
|
@ -23,44 +164,64 @@ If you use Ollama, you can use any pulled model that supports embeddings.
|
||||||
|
|
||||||
### Ollama (Default)
|
### Ollama (Default)
|
||||||
|
|
||||||
```bash
|
```yaml
|
||||||
EMBEDDINGS_PROVIDER="ollama"
|
embeddings:
|
||||||
EMBEDDINGS_MODEL="mxbai-embed-large"
|
provider: ollama
|
||||||
EMBEDDINGS_VECTOR_DIM=1024
|
model: mxbai-embed-large
|
||||||
|
vector_dim: 1024
|
||||||
```
|
```
|
||||||
|
|
||||||
### VoyageAI
|
### VoyageAI
|
||||||
If you want to use VoyageAI embeddings you will need to install `haiku.rag` with the VoyageAI extras,
|
|
||||||
|
If you want to use VoyageAI embeddings you will need to install `haiku.rag` with the VoyageAI extras:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
uv pip install haiku.rag[voyageai]
|
uv pip install haiku.rag[voyageai]
|
||||||
```
|
```
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
embeddings:
|
||||||
|
provider: voyageai
|
||||||
|
model: voyage-3.5
|
||||||
|
vector_dim: 1024
|
||||||
|
```
|
||||||
|
|
||||||
|
Set your API key via environment variable:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
EMBEDDINGS_PROVIDER="voyageai"
|
export VOYAGE_API_KEY=your-api-key
|
||||||
EMBEDDINGS_MODEL="voyage-3.5"
|
|
||||||
EMBEDDINGS_VECTOR_DIM=1024
|
|
||||||
VOYAGE_API_KEY="your-api-key"
|
|
||||||
```
|
```
|
||||||
|
|
||||||
### OpenAI
|
### OpenAI
|
||||||
OpenAI embeddings are included in the default installation. Simply set environment variables:
|
|
||||||
|
OpenAI embeddings are included in the default installation:
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
embeddings:
|
||||||
|
provider: openai
|
||||||
|
model: text-embedding-3-small # or text-embedding-3-large
|
||||||
|
vector_dim: 1536
|
||||||
|
```
|
||||||
|
|
||||||
|
Set your API key via environment variable:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
EMBEDDINGS_PROVIDER="openai"
|
export OPENAI_API_KEY=your-api-key
|
||||||
EMBEDDINGS_MODEL="text-embedding-3-small" # or text-embedding-3-large
|
|
||||||
EMBEDDINGS_VECTOR_DIM=1536
|
|
||||||
OPENAI_API_KEY="your-api-key"
|
|
||||||
```
|
```
|
||||||
|
|
||||||
### vLLM
|
### vLLM
|
||||||
|
|
||||||
For high-performance local inference, you can use vLLM to serve embedding models with OpenAI-compatible APIs:
|
For high-performance local inference, you can use vLLM to serve embedding models with OpenAI-compatible APIs:
|
||||||
|
|
||||||
```bash
|
```yaml
|
||||||
EMBEDDINGS_PROVIDER="vllm"
|
embeddings:
|
||||||
EMBEDDINGS_MODEL="mixedbread-ai/mxbai-embed-large-v1" # Any embedding model supported by vLLM
|
provider: vllm
|
||||||
EMBEDDINGS_VECTOR_DIM=512 # Dimension depends on the model
|
model: mixedbread-ai/mxbai-embed-large-v1
|
||||||
VLLM_EMBEDDINGS_BASE_URL="http://localhost:8000" # vLLM server URL
|
vector_dim: 512
|
||||||
|
|
||||||
|
providers:
|
||||||
|
vllm:
|
||||||
|
embeddings_base_url: http://localhost:8000
|
||||||
```
|
```
|
||||||
|
|
||||||
**Note:** You need to run a vLLM server separately with an embedding model loaded.
|
**Note:** You need to run a vLLM server separately with an embedding model loaded.
|
||||||
|
|
@ -71,60 +232,83 @@ Configure which LLM provider to use for question answering. Any provider and mod
|
||||||
|
|
||||||
### Ollama (Default)
|
### Ollama (Default)
|
||||||
|
|
||||||
```bash
|
```yaml
|
||||||
QA_PROVIDER="ollama"
|
qa:
|
||||||
QA_MODEL="gpt-oss"
|
provider: ollama
|
||||||
OLLAMA_BASE_URL="http://localhost:11434"
|
model: gpt-oss
|
||||||
|
|
||||||
|
providers:
|
||||||
|
ollama:
|
||||||
|
base_url: http://localhost:11434
|
||||||
```
|
```
|
||||||
|
|
||||||
### OpenAI
|
### OpenAI
|
||||||
|
|
||||||
OpenAI QA is included in the default installation. Simply configure:
|
OpenAI QA is included in the default installation:
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
qa:
|
||||||
|
provider: openai
|
||||||
|
model: gpt-4o-mini # or gpt-4, gpt-3.5-turbo, etc.
|
||||||
|
```
|
||||||
|
|
||||||
|
Set your API key via environment variable:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
QA_PROVIDER="openai"
|
export OPENAI_API_KEY=your-api-key
|
||||||
QA_MODEL="gpt-4o-mini" # or gpt-4, gpt-3.5-turbo, etc.
|
|
||||||
OPENAI_API_KEY="your-api-key"
|
|
||||||
```
|
```
|
||||||
|
|
||||||
### Anthropic
|
### Anthropic
|
||||||
|
|
||||||
Anthropic QA is included in the default installation. Simply configure:
|
Anthropic QA is included in the default installation:
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
qa:
|
||||||
|
provider: anthropic
|
||||||
|
model: claude-3-5-haiku-20241022 # or claude-3-5-sonnet-20241022, etc.
|
||||||
|
```
|
||||||
|
|
||||||
|
Set your API key via environment variable:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
QA_PROVIDER="anthropic"
|
export ANTHROPIC_API_KEY=your-api-key
|
||||||
QA_MODEL="claude-3-5-haiku-20241022" # or claude-3-5-sonnet-20241022, etc.
|
|
||||||
ANTHROPIC_API_KEY="your-api-key"
|
|
||||||
```
|
```
|
||||||
|
|
||||||
### vLLM
|
### vLLM
|
||||||
|
|
||||||
For high-performance local inference, you can use vLLM to serve models with OpenAI-compatible APIs:
|
For high-performance local inference:
|
||||||
|
|
||||||
```bash
|
```yaml
|
||||||
QA_PROVIDER="vllm"
|
qa:
|
||||||
QA_MODEL="Qwen/Qwen3-4B" # Any model with tool support in vLLM
|
provider: vllm
|
||||||
VLLM_QA_BASE_URL="http://localhost:8002" # vLLM server URL
|
model: Qwen/Qwen3-4B # Any model with tool support in vLLM
|
||||||
|
|
||||||
|
providers:
|
||||||
|
vllm:
|
||||||
|
qa_base_url: http://localhost:8002
|
||||||
```
|
```
|
||||||
|
|
||||||
**Note:** You need to run a vLLM server separately with a model that supports tool calling loaded. Consult the specific model's documentation for proper vLLM serving configuration.
|
**Note:** You need to run a vLLM server separately with a model that supports tool calling loaded. Consult the specific model's documentation for proper vLLM serving configuration.
|
||||||
|
|
||||||
### Other Providers
|
### Other Providers
|
||||||
|
|
||||||
Any provider supported by Pydantic AI can be used. Examples include:
|
Any provider supported by Pydantic AI can be used. Examples:
|
||||||
|
|
||||||
```bash
|
```yaml
|
||||||
# Google Gemini
|
# Google Gemini
|
||||||
QA_PROVIDER="gemini"
|
qa:
|
||||||
QA_MODEL="gemini-1.5-flash"
|
provider: gemini
|
||||||
|
model: gemini-1.5-flash
|
||||||
|
|
||||||
# Groq
|
# Groq
|
||||||
QA_PROVIDER="groq"
|
qa:
|
||||||
QA_MODEL="llama-3.3-70b-versatile"
|
provider: groq
|
||||||
|
model: llama-3.3-70b-versatile
|
||||||
|
|
||||||
# Mistral
|
# Mistral
|
||||||
QA_PROVIDER="mistral"
|
qa:
|
||||||
QA_MODEL="mistral-small-latest"
|
provider: mistral
|
||||||
|
model: mistral-small-latest
|
||||||
```
|
```
|
||||||
|
|
||||||
See the [Pydantic AI documentation](https://ai.pydantic.dev/models/) for the complete list of supported providers and models.
|
See the [Pydantic AI documentation](https://ai.pydantic.dev/models/) for the complete list of supported providers and models.
|
||||||
|
|
@ -133,7 +317,7 @@ See the [Pydantic AI documentation](https://ai.pydantic.dev/models/) for the com
|
||||||
|
|
||||||
Reranking improves search quality by re-ordering the initial search results using specialized models. When enabled, the system retrieves more candidates (3x the requested limit) and then reranks them to return the most relevant results.
|
Reranking improves search quality by re-ordering the initial search results using specialized models. When enabled, the system retrieves more candidates (3x the requested limit) and then reranks them to return the most relevant results.
|
||||||
|
|
||||||
Reranking is **disabled by default** (`RERANK_PROVIDER=""`) for faster searches. You can enable it by configuring one of the providers below.
|
Reranking is **disabled by default** (`provider: ""`) for faster searches. You can enable it by configuring one of the providers below.
|
||||||
|
|
||||||
### MixedBread AI
|
### MixedBread AI
|
||||||
|
|
||||||
|
|
@ -145,29 +329,40 @@ uv pip install haiku.rag[mxbai]
|
||||||
|
|
||||||
Then configure:
|
Then configure:
|
||||||
|
|
||||||
```bash
|
```yaml
|
||||||
RERANK_PROVIDER="mxbai"
|
reranking:
|
||||||
RERANK_MODEL="mixedbread-ai/mxbai-rerank-base-v2"
|
provider: mxbai
|
||||||
|
model: mixedbread-ai/mxbai-rerank-base-v2
|
||||||
```
|
```
|
||||||
|
|
||||||
### Cohere
|
### Cohere
|
||||||
|
|
||||||
Cohere reranking is included in the default installation. Simply configure:
|
Cohere reranking is included in the default installation:
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
reranking:
|
||||||
|
provider: cohere
|
||||||
|
model: rerank-v3.5
|
||||||
|
```
|
||||||
|
|
||||||
|
Set your API key via environment variable:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
RERANK_PROVIDER="cohere"
|
export CO_API_KEY=your-api-key
|
||||||
RERANK_MODEL="rerank-v3.5"
|
|
||||||
COHERE_API_KEY="your-api-key"
|
|
||||||
```
|
```
|
||||||
|
|
||||||
### vLLM
|
### vLLM
|
||||||
|
|
||||||
For high-performance local reranking using dedicated reranking models:
|
For high-performance local reranking using dedicated reranking models:
|
||||||
|
|
||||||
```bash
|
```yaml
|
||||||
RERANK_PROVIDER="vllm"
|
reranking:
|
||||||
RERANK_MODEL="mixedbread-ai/mxbai-rerank-base-v2" # Any reranking model supported by vLLM
|
provider: vllm
|
||||||
VLLM_RERANK_BASE_URL="http://localhost:8001" # vLLM server URL
|
model: mixedbread-ai/mxbai-rerank-base-v2
|
||||||
|
|
||||||
|
providers:
|
||||||
|
vllm:
|
||||||
|
rerank_base_url: http://localhost:8001
|
||||||
```
|
```
|
||||||
|
|
||||||
**Note:** vLLM reranking uses the `/rerank` API endpoint. You need to run a vLLM server separately with a reranking model loaded. Consult the specific model's documentation for proper vLLM serving configuration.
|
**Note:** vLLM reranking uses the `/rerank` API endpoint. You need to run a vLLM server separately with a reranking model loaded. Consult the specific model's documentation for proper vLLM serving configuration.
|
||||||
|
|
@ -178,78 +373,91 @@ VLLM_RERANK_BASE_URL="http://localhost:8001" # vLLM server URL
|
||||||
|
|
||||||
By default, `haiku.rag` uses a local LanceDB database:
|
By default, `haiku.rag` uses a local LanceDB database:
|
||||||
|
|
||||||
```bash
|
```yaml
|
||||||
# Default data directory (where local LanceDB is stored)
|
storage:
|
||||||
DEFAULT_DATA_DIR="/path/to/data"
|
data_dir: /path/to/data # Empty = use default platform location
|
||||||
```
|
```
|
||||||
|
|
||||||
For remote storage, use the `LANCEDB_URI` setting with various backends:
|
For remote storage, use the `lancedb` settings with various backends:
|
||||||
|
|
||||||
```bash
|
```yaml
|
||||||
# LanceDB Cloud
|
# LanceDB Cloud
|
||||||
LANCEDB_URI="db://your-database-name"
|
lancedb:
|
||||||
LANCEDB_API_KEY="your-api-key"
|
uri: db://your-database-name
|
||||||
LANCEDB_REGION="us-west-2" # optional
|
api_key: your-api-key
|
||||||
|
region: us-west-2 # optional
|
||||||
|
|
||||||
# Amazon S3
|
# Amazon S3
|
||||||
LANCEDB_URI="s3://my-bucket/my-table"
|
lancedb:
|
||||||
|
uri: s3://my-bucket/my-table
|
||||||
# Use AWS credentials or IAM roles
|
# Use AWS credentials or IAM roles
|
||||||
|
|
||||||
# Azure Blob Storage
|
# Azure Blob Storage
|
||||||
LANCEDB_URI="az://my-container/my-table"
|
lancedb:
|
||||||
|
uri: az://my-container/my-table
|
||||||
# Use Azure credentials
|
# Use Azure credentials
|
||||||
|
|
||||||
# Google Cloud Storage
|
# Google Cloud Storage
|
||||||
LANCEDB_URI="gs://my-bucket/my-table"
|
lancedb:
|
||||||
|
uri: gs://my-bucket/my-table
|
||||||
# Use GCP credentials
|
# Use GCP credentials
|
||||||
|
|
||||||
# HDFS
|
# HDFS
|
||||||
LANCEDB_URI="hdfs://namenode:port/path/to/table"
|
lancedb:
|
||||||
|
uri: hdfs://namenode:port/path/to/table
|
||||||
```
|
```
|
||||||
|
|
||||||
Authentication is handled through standard cloud provider credentials (AWS CLI, Azure CLI, gcloud, etc.) or by setting `LANCEDB_API_KEY` for LanceDB Cloud.
|
Authentication is handled through standard cloud provider credentials (AWS CLI, Azure CLI, gcloud, etc.) or by setting `api_key` for LanceDB Cloud.
|
||||||
|
|
||||||
**Note:** Table optimization is automatically handled by LanceDB Cloud (`db://` URIs) and is disabled for better performance. For object storage backends (S3, Azure, GCS), optimization is still performed locally.
|
**Note:** Table optimization is automatically handled by LanceDB Cloud (`db://` URIs) and is disabled for better performance. For object storage backends (S3, Azure, GCS), optimization is still performed locally.
|
||||||
|
|
||||||
#### Disable database auto-creation
|
#### Disable database auto-creation
|
||||||
|
|
||||||
By default, haiku.rag creates the local LanceDB directory and required tables on first use. To prevent accidental database creation and fail fast if a database hasn’t been set up yet, set:
|
By default, haiku.rag creates the local LanceDB directory and required tables on first use. To prevent accidental database creation and fail fast if a database hasn't been set up yet:
|
||||||
|
|
||||||
```bash
|
```yaml
|
||||||
DISABLE_DB_AUTOCREATE=true
|
storage:
|
||||||
|
disable_autocreate: true
|
||||||
```
|
```
|
||||||
|
|
||||||
When enabled, for local paths, haiku.rag errors if the LanceDB directory does not exist, and it will not create parent directories.
|
When enabled, for local paths, haiku.rag errors if the LanceDB directory does not exist, and it will not create parent directories.
|
||||||
|
|
||||||
### Document Processing
|
### Document Processing
|
||||||
|
|
||||||
```bash
|
```yaml
|
||||||
# Chunk size for document processing
|
processing:
|
||||||
CHUNK_SIZE=256
|
# Chunk size for document processing
|
||||||
|
chunk_size: 256
|
||||||
|
|
||||||
# Number of adjacent chunks to include before/after retrieved chunks for context
|
# Number of adjacent chunks to include before/after retrieved chunks for context
|
||||||
# 0 = no expansion (default), 1 = include 1 chunk before and after, etc.
|
# 0 = no expansion (default), 1 = include 1 chunk before and after, etc.
|
||||||
# When expanded chunks overlap or are adjacent, they are automatically merged
|
# When expanded chunks overlap or are adjacent, they are automatically merged
|
||||||
# into single chunks with continuous content to eliminate duplication
|
# into single chunks with continuous content to eliminate duplication
|
||||||
CONTEXT_CHUNK_RADIUS=0
|
context_chunk_radius: 0
|
||||||
|
|
||||||
# Vacuum retention threshold (seconds) for automatic cleanup
|
# Optional dotted path or file path to a callable that preprocesses
|
||||||
# When documents are added/updated, old table versions older than this are removed
|
# markdown content before chunking
|
||||||
# Default: 60 seconds (safe for concurrent connections)
|
markdown_preprocessor: ""
|
||||||
# Set to 0 for aggressive cleanup (removes all old versions immediately)
|
|
||||||
VACUUM_RETENTION_SECONDS=60
|
storage:
|
||||||
|
# Vacuum retention threshold (seconds) for automatic cleanup
|
||||||
|
# When documents are added/updated, old table versions older than this are removed
|
||||||
|
# Default: 60 seconds (safe for concurrent connections)
|
||||||
|
# Set to 0 for aggressive cleanup (removes all old versions immediately)
|
||||||
|
vacuum_retention_seconds: 60
|
||||||
```
|
```
|
||||||
|
|
||||||
#### Markdown Preprocessor
|
#### Markdown Preprocessor
|
||||||
|
|
||||||
Optionally preprocess Markdown before chunking by pointing to a callable that receives and returns Markdown text. This is useful for normalizing content, stripping boilerplate, or applying custom transformations before chunk boundaries are computed.
|
Optionally preprocess Markdown before chunking by pointing to a callable that receives and returns Markdown text. This is useful for normalizing content, stripping boilerplate, or applying custom transformations before chunk boundaries are computed.
|
||||||
|
|
||||||
```bash
|
```yaml
|
||||||
# A callable path in one of these formats:
|
processing:
|
||||||
# - package.module:func
|
# A callable path in one of these formats:
|
||||||
# - package.module.func
|
# - package.module:func
|
||||||
# - /abs/or/relative/path/to/file.py:func
|
# - package.module.func
|
||||||
MARKDOWN_PREPROCESSOR="my_pkg.preprocess:clean_md"
|
# - /abs/or/relative/path/to/file.py:func
|
||||||
|
markdown_preprocessor: my_pkg.preprocess:clean_md
|
||||||
```
|
```
|
||||||
|
|
||||||
!!! note
|
!!! note
|
||||||
|
|
@ -271,3 +479,26 @@ def clean_md(text: str) -> str:
|
||||||
out.append(line)
|
out.append(line)
|
||||||
return "\n".join(out)
|
return "\n".join(out)
|
||||||
```
|
```
|
||||||
|
|
||||||
|
## Migration from Environment Variables
|
||||||
|
|
||||||
|
!!! warning "Deprecation Notice"
|
||||||
|
Environment variable configuration via `.env` files is deprecated and will be removed in future versions. Please migrate to YAML configuration.
|
||||||
|
|
||||||
|
To migrate your existing `.env` file to YAML:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
haiku-rag init-config --from-env
|
||||||
|
```
|
||||||
|
|
||||||
|
This will read your current environment variables and generate a `haiku.rag.yaml` file with those settings.
|
||||||
|
|
||||||
|
!!! note
|
||||||
|
When migrating from environment variables, list values like `MONITOR_DIRECTORIES` that were comma-separated (`/path1,/path2`) will be converted to proper YAML lists. In YAML, always use list syntax:
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
storage:
|
||||||
|
monitor_directories:
|
||||||
|
- /path/to/dir1
|
||||||
|
- /path/to/dir2
|
||||||
|
```
|
||||||
|
|
|
||||||
|
|
@ -47,26 +47,31 @@ vllm serve Qwen/Qwen3-4B --port 8002 --enable-auto-tool-choice --tool-call-parse
|
||||||
vllm serve mixedbread-ai/mxbai-rerank-base-v2 --hf_overrides '{"architectures": ["Qwen2ForSequenceClassification"],"classifier_from_token": ["0", "1"], "method": "from_2_way_softmax"}' --port 8001
|
vllm serve mixedbread-ai/mxbai-rerank-base-v2 --hf_overrides '{"architectures": ["Qwen2ForSequenceClassification"],"classifier_from_token": ["0", "1"], "method": "from_2_way_softmax"}' --port 8001
|
||||||
```
|
```
|
||||||
|
|
||||||
Then configure haiku.rag to use the vLLM servers:
|
Then configure haiku.rag to use the vLLM servers. Create a `haiku.rag.yaml` file:
|
||||||
|
|
||||||
```bash
|
```yaml
|
||||||
# Embeddings
|
embeddings:
|
||||||
EMBEDDINGS_PROVIDER="vllm"
|
provider: vllm
|
||||||
EMBEDDINGS_MODEL="mixedbread-ai/mxbai-embed-large-v1"
|
model: mixedbread-ai/mxbai-embed-large-v1
|
||||||
EMBEDDINGS_VECTOR_DIM=512
|
vector_dim: 512
|
||||||
VLLM_EMBEDDINGS_BASE_URL="http://localhost:8000"
|
|
||||||
|
|
||||||
# QA (optional)
|
qa:
|
||||||
QA_PROVIDER="vllm"
|
provider: vllm
|
||||||
QA_MODEL="Qwen/Qwen3-4B"
|
model: Qwen/Qwen3-4B
|
||||||
VLLM_QA_BASE_URL="http://localhost:8002"
|
|
||||||
|
|
||||||
# Reranking (optional)
|
reranking:
|
||||||
RERANK_PROVIDER="vllm"
|
provider: vllm
|
||||||
RERANK_MODEL="mixedbread-ai/mxbai-rerank-base-v2"
|
model: mixedbread-ai/mxbai-rerank-base-v2
|
||||||
VLLM_RERANK_BASE_URL="http://localhost:8001"
|
|
||||||
|
providers:
|
||||||
|
vllm:
|
||||||
|
embeddings_base_url: http://localhost:8000
|
||||||
|
qa_base_url: http://localhost:8002
|
||||||
|
rerank_base_url: http://localhost:8001
|
||||||
```
|
```
|
||||||
|
|
||||||
|
See [Configuration](configuration.md) for all available options.
|
||||||
|
|
||||||
## Requirements
|
## Requirements
|
||||||
|
|
||||||
- Python 3.12+
|
- Python 3.12+
|
||||||
|
|
|
||||||
|
|
@ -189,7 +189,7 @@ for chunk, score in expanded_results:
|
||||||
|
|
||||||
**Smart Merging**: When expanded chunks overlap or are adjacent within the same document, they are automatically merged into single chunks with continuous content. This eliminates duplication and provides coherent text blocks. The merged chunk uses the highest relevance score from the original chunks.
|
**Smart Merging**: When expanded chunks overlap or are adjacent within the same document, they are automatically merged into single chunks with continuous content. This eliminates duplication and provides coherent text blocks. The merged chunk uses the highest relevance score from the original chunks.
|
||||||
|
|
||||||
This is automatically used by the QA system when `CONTEXT_CHUNK_RADIUS > 0` to provide better answers with more complete context.
|
This is automatically used by the QA system when `processing.context_chunk_radius > 0` (configured in `haiku.rag.yaml`) to provide better answers with more complete context.
|
||||||
|
|
||||||
## Question Answering
|
## Question Answering
|
||||||
|
|
||||||
|
|
@ -222,6 +222,6 @@ answer = await client.ask(
|
||||||
|
|
||||||
The QA agent will search your documents for relevant information and use the configured LLM to generate a comprehensive answer. With `cite=True`, responses include citations showing which documents were used as sources. Citations prefer the document title when present, otherwise they use the URI.
|
The QA agent will search your documents for relevant information and use the configured LLM to generate a comprehensive answer. With `cite=True`, responses include citations showing which documents were used as sources. Citations prefer the document title when present, otherwise they use the URI.
|
||||||
|
|
||||||
The QA provider and model can be configured via environment variables (see [Configuration](configuration.md)).
|
The QA provider and model are configured in `haiku.rag.yaml` or can be passed directly to the client (see [Configuration](configuration.md)).
|
||||||
|
|
||||||
See also: [Agents](agents.md) for details on the QA agent and the multi‑agent research workflow.
|
See also: [Agents](agents.md) for details on the QA agent and the multi‑agent research workflow.
|
||||||
|
|
|
||||||
|
|
@ -45,11 +45,19 @@ This will start file monitoring, MCP server on port 8001, and A2A server on port
|
||||||
|
|
||||||
## File Monitoring
|
## File Monitoring
|
||||||
|
|
||||||
Set `MONITOR_DIRECTORIES` environment variable to enable automatic file monitoring:
|
Configure directories to monitor in your `haiku.rag.yaml`:
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
storage:
|
||||||
|
monitor_directories:
|
||||||
|
- /path/to/documents
|
||||||
|
- /another/path
|
||||||
|
```
|
||||||
|
|
||||||
|
Then start the server:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
export MONITOR_DIRECTORIES="/path/to/documents"
|
haiku-rag serve --monitor
|
||||||
haiku-rag serve
|
|
||||||
```
|
```
|
||||||
|
|
||||||
### Monitoring Features
|
### Monitoring Features
|
||||||
|
|
|
||||||
|
|
@ -31,34 +31,28 @@ Install `haiku.rag` Python package using [uv](https://docs.astral.sh/uv/getting-
|
||||||
uv pip install haiku.rag
|
uv pip install haiku.rag
|
||||||
```
|
```
|
||||||
|
|
||||||
Configure your OpenAI API key and embeddings model.
|
Configure haiku.rag to use OpenAI. Create a `haiku.rag.yaml` file:
|
||||||
|
|
||||||
- Haiku RAG supports [dotenv](https://pypi.org/project/python-dotenv/) environment files and environment varibles for configuration
|
```yaml
|
||||||
- [See OpenAPI vector embeddings documentation](https://platform.openai.com/docs/guides/embeddings/embedding-models)
|
embeddings:
|
||||||
|
provider: openai
|
||||||
|
model: text-embedding-3-small # or text-embedding-3-large
|
||||||
|
vector_dim: 1536
|
||||||
|
|
||||||
Create a file called `.env` and add:
|
qa:
|
||||||
|
provider: openai
|
||||||
```shell
|
model: gpt-4o-mini # or gpt-4o, gpt-4, etc.
|
||||||
#
|
|
||||||
# These settings are relevant for converting documents to embeddings
|
|
||||||
#
|
|
||||||
|
|
||||||
EMBEDDINGS_PROVIDER="openai"
|
|
||||||
# or text-embedding-3-large
|
|
||||||
EMBEDDINGS_MODEL="text-embedding-3-small"
|
|
||||||
EMBEDDINGS_VECTOR_DIM=1536
|
|
||||||
OPENAI_API_KEY="<your OpenAPI API key goes here>"
|
|
||||||
|
|
||||||
#
|
|
||||||
# These settings are relevant for question answering chats
|
|
||||||
#
|
|
||||||
|
|
||||||
# We tell Haiku.rag to use OpenAI remote AI for chats, instead of local ollama.
|
|
||||||
QA_PROVIDER="openai"
|
|
||||||
QA_MODEL="gpt-4o-mini" # or gpt-4, gpt-3.5-turbo, etc.
|
|
||||||
```
|
```
|
||||||
|
|
||||||
For the list of available OpenAI embedding models and `EMBEDDINGS_VECTOR_DIM` options, ask ChatGPT for instructions.
|
Set your OpenAI API key as an environment variable (API keys should not be stored in the YAML file):
|
||||||
|
|
||||||
|
```bash
|
||||||
|
export OPENAI_API_KEY="<your OpenAI API key>"
|
||||||
|
```
|
||||||
|
|
||||||
|
For the list of available OpenAI models and their vector dimensions, see the [OpenAI documentation](https://platform.openai.com/docs/guides/embeddings).
|
||||||
|
|
||||||
|
See [Configuration](configuration.md) for all available options.
|
||||||
|
|
||||||
## Adding the first documents
|
## Adding the first documents
|
||||||
|
|
||||||
|
|
@ -231,6 +225,4 @@ rm -rf "/Users/moo/Library/Application Support/haiku.rag/haiku.rag.lancedb"
|
||||||
|
|
||||||
## Configuration
|
## Configuration
|
||||||
|
|
||||||
See [Configuration page](./configuration.md) for more information about configuration
|
See [Configuration page](./configuration.md) for complete documentation on YAML configuration and all available options.
|
||||||
|
|
||||||
For the available environment variable config options see [config.py](https://github.com/ggozad/haiku.rag/blob/main/src/haiku/rag/config.py).
|
|
||||||
|
|
|
||||||
|
|
@ -1,14 +1,3 @@
|
||||||
# QA Provider for the research agent (ollama, openai, anthropic, etc.)
|
|
||||||
QA_PROVIDER=ollama
|
|
||||||
|
|
||||||
# QA Model name
|
|
||||||
QA_MODEL=gpt-oss:latest
|
|
||||||
|
|
||||||
# Ollama base URL (only needed if using ollama provider)
|
|
||||||
# For Docker: http://host.docker.internal:11434
|
|
||||||
# For local development: http://localhost:11434
|
|
||||||
OLLAMA_BASE_URL=http://host.docker.internal:11434
|
|
||||||
|
|
||||||
# Path to the LanceDB database
|
# Path to the LanceDB database
|
||||||
# For Docker: /app/data/haiku_rag.lancedb
|
# For Docker: /app/data/haiku_rag.lancedb
|
||||||
# For local development: Use absolute path to existing database
|
# For local development: Use absolute path to existing database
|
||||||
|
|
@ -17,7 +6,5 @@ DB_PATH=~/SOME_FOLDER/haiku.rag.lancedb
|
||||||
# API keys (set as needed for your QA provider)
|
# API keys (set as needed for your QA provider)
|
||||||
# OPENAI_API_KEY=your-key-here
|
# OPENAI_API_KEY=your-key-here
|
||||||
# ANTHROPIC_API_KEY=your-key-here
|
# ANTHROPIC_API_KEY=your-key-here
|
||||||
|
# VOYAGE_API_KEY=your-key-here
|
||||||
# Embedding provider configuration (optional, defaults will be used)
|
# CO_API_KEY=your-key-here
|
||||||
# EMBEDDING_PROVIDER=openai
|
|
||||||
# EMBEDDING_MODEL=text-embedding-3-small
|
|
||||||
|
|
|
||||||
3
examples/ag-ui-research/.gitignore
vendored
Normal file
3
examples/ag-ui-research/.gitignore
vendored
Normal file
|
|
@ -0,0 +1,3 @@
|
||||||
|
haiku.rag.yaml
|
||||||
|
.env
|
||||||
|
data/
|
||||||
|
|
@ -30,19 +30,27 @@ Research assistant powered by [haiku.rag](https://ggozad.github.io/haiku.rag/),
|
||||||
haiku-rag add-src document.pdf --db data/haiku_rag.lancedb
|
haiku-rag add-src document.pdf --db data/haiku_rag.lancedb
|
||||||
```
|
```
|
||||||
|
|
||||||
2. **Configure environment** (optional)
|
2. **Configure haiku.rag**
|
||||||
```bash
|
```bash
|
||||||
cp .env.example .env
|
cp haiku.rag.yaml.example haiku.rag.yaml
|
||||||
# Edit .env to customize provider/model
|
# Edit haiku.rag.yaml to customize provider/model
|
||||||
```
|
```
|
||||||
See [haiku.rag configuration](https://ggozad.github.io/haiku.rag/configuration/) for details.
|
See [haiku.rag configuration](https://ggozad.github.io/haiku.rag/configuration/) for details.
|
||||||
|
|
||||||
3. **Start the application**
|
3. **Set API keys** (if using non-Ollama providers)
|
||||||
|
```bash
|
||||||
|
cp .env.example .env
|
||||||
|
# Edit .env to set your API keys
|
||||||
|
export OPENAI_API_KEY=your-key-here
|
||||||
|
export ANTHROPIC_API_KEY=your-key-here
|
||||||
|
```
|
||||||
|
|
||||||
|
4. **Start the application**
|
||||||
```bash
|
```bash
|
||||||
docker compose up --build
|
docker compose up --build
|
||||||
```
|
```
|
||||||
|
|
||||||
4. **Access the interface**
|
5. **Access the interface**
|
||||||
- Frontend: http://localhost:3000
|
- Frontend: http://localhost:3000
|
||||||
- Backend health: http://localhost:8000/health
|
- Backend health: http://localhost:8000/health
|
||||||
|
|
||||||
|
|
@ -60,6 +68,7 @@ Research assistant powered by [haiku.rag](https://ggozad.github.io/haiku.rag/),
|
||||||
## Architecture
|
## Architecture
|
||||||
|
|
||||||
- **Backend** (Python): Pydantic AI agent with haiku.rag integration
|
- **Backend** (Python): Pydantic AI agent with haiku.rag integration
|
||||||
|
- Uses published `ghcr.io/ggozad/haiku.rag:latest` Docker image as base
|
||||||
- `agent.py`: Research agent with tool definitions
|
- `agent.py`: Research agent with tool definitions
|
||||||
- `main.py`: Starlette app serving AG-UI protocol
|
- `main.py`: Starlette app serving AG-UI protocol
|
||||||
|
|
||||||
|
|
@ -70,11 +79,15 @@ Research assistant powered by [haiku.rag](https://ggozad.github.io/haiku.rag/),
|
||||||
|
|
||||||
## Configuration
|
## Configuration
|
||||||
|
|
||||||
|
Configuration is done through `haiku.rag.yaml` (see `haiku.rag.yaml.example`):
|
||||||
|
|
||||||
|
- `qa.provider`: LLM provider (default: `ollama`)
|
||||||
|
- `qa.model`: Model name (default: `gpt-oss:latest`)
|
||||||
|
- `providers.ollama.base_url`: Ollama endpoint (default: `http://host.docker.internal:11434`)
|
||||||
|
|
||||||
Environment variables (see `.env.example`):
|
Environment variables (see `.env.example`):
|
||||||
|
|
||||||
- `DB_PATH`: Path to haiku.rag database (default: `haiku_rag.lancedb`)
|
- `DB_PATH`: Path to haiku.rag database (default: `haiku_rag.lancedb`)
|
||||||
- `QA_PROVIDER`: LLM provider (default: `ollama`)
|
- `OPENAI_API_KEY`, `ANTHROPIC_API_KEY`: API keys for cloud providers
|
||||||
- `QA_MODEL`: Model name (default: `gpt-oss:latest`)
|
|
||||||
- `OLLAMA_BASE_URL`: Ollama endpoint (default: `http://host.docker.internal:11434`)
|
|
||||||
|
|
||||||
For other providers (OpenAI, Anthropic, etc.), see [haiku.rag configuration docs](https://ggozad.github.io/haiku.rag/configuration/).
|
For full configuration options, see [haiku.rag configuration docs](https://ggozad.github.io/haiku.rag/configuration/).
|
||||||
|
|
|
||||||
|
|
@ -1,27 +1,15 @@
|
||||||
FROM ghcr.io/astral-sh/uv:python3.13-bookworm-slim
|
FROM ghcr.io/ggozad/haiku.rag:latest
|
||||||
|
|
||||||
WORKDIR /app
|
WORKDIR /app
|
||||||
|
|
||||||
# Enable bytecode compilation
|
# Copy backend application files
|
||||||
ENV UV_COMPILE_BYTECODE=1
|
COPY agent.py main.py ./
|
||||||
|
COPY pyproject.toml uv.lock ./
|
||||||
|
|
||||||
# Copy from the cache instead of linking since it's a mounted volume
|
# Install backend dependencies into the existing haiku.rag venv
|
||||||
ENV UV_LINK_MODE=copy
|
RUN uv sync --frozen --no-dev
|
||||||
|
|
||||||
# Install dependencies
|
|
||||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
|
||||||
--mount=type=bind,source=uv.lock,target=uv.lock \
|
|
||||||
--mount=type=bind,source=pyproject.toml,target=pyproject.toml \
|
|
||||||
uv sync --frozen --no-install-project --no-dev
|
|
||||||
|
|
||||||
# Copy the project into the image
|
|
||||||
COPY . .
|
|
||||||
|
|
||||||
# Sync the project
|
|
||||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
|
||||||
uv sync --frozen --no-dev
|
|
||||||
|
|
||||||
EXPOSE 8000
|
EXPOSE 8000
|
||||||
|
|
||||||
# Run with uv
|
# Run with uvicorn
|
||||||
CMD ["uv", "run", "uvicorn", "main:app", "--host", "0.0.0.0", "--port", "8000", "--reload"]
|
CMD ["uv", "run", "uvicorn", "main:app", "--host", "0.0.0.0", "--port", "8000", "--reload"]
|
||||||
|
|
|
||||||
|
|
@ -38,13 +38,13 @@ def _as_state_snapshot(ctx: RunContext[ResearchDeps]) -> StateSnapshotEvent:
|
||||||
|
|
||||||
|
|
||||||
def create_agent(
|
def create_agent(
|
||||||
qa_provider: str = Config.QA_PROVIDER, qa_model: str = Config.QA_MODEL
|
qa_provider: str = Config.qa.provider, qa_model: str = Config.qa.model
|
||||||
) -> Agent[ResearchDeps, str]:
|
) -> Agent[ResearchDeps, str]:
|
||||||
"""Create and configure the research agent.
|
"""Create and configure the research agent.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
qa_provider: QA provider for the agent (default: from Config.QA_PROVIDER)
|
qa_provider: QA provider for the agent (default: from Config.qa.provider)
|
||||||
qa_model: Model name to use (default: from Config.QA_MODEL)
|
qa_model: Model name to use (default: from Config.qa.model)
|
||||||
"""
|
"""
|
||||||
print(f"[AGENT SETUP] Creating agent with provider={qa_provider}, model={qa_model}")
|
print(f"[AGENT SETUP] Creating agent with provider={qa_provider}, model={qa_model}")
|
||||||
agent = Agent(
|
agent = Agent(
|
||||||
|
|
|
||||||
|
|
@ -33,7 +33,7 @@ async def lifespan(app):
|
||||||
logger.info(f"Initializing HaikuRAG client with database: {db_path}")
|
logger.info(f"Initializing HaikuRAG client with database: {db_path}")
|
||||||
client = HaikuRAG(db_path)
|
client = HaikuRAG(db_path)
|
||||||
logger.info("Research assistant backend ready")
|
logger.info("Research assistant backend ready")
|
||||||
logger.info(f"QA Provider: {Config.QA_PROVIDER}, Model: {Config.QA_MODEL}")
|
logger.info(f"QA Provider: {Config.qa.provider}, Model: {Config.qa.model}")
|
||||||
|
|
||||||
yield
|
yield
|
||||||
|
|
||||||
|
|
@ -51,9 +51,9 @@ async def health(request):
|
||||||
{
|
{
|
||||||
"status": "healthy",
|
"status": "healthy",
|
||||||
"agent_model": str(agent.model),
|
"agent_model": str(agent.model),
|
||||||
"qa_provider": Config.QA_PROVIDER,
|
"qa_provider": Config.qa.provider,
|
||||||
"qa_model": Config.QA_MODEL,
|
"qa_model": Config.qa.model,
|
||||||
"ollama_base_url": Config.OLLAMA_BASE_URL,
|
"ollama_base_url": Config.providers.ollama.base_url,
|
||||||
"db_path": db_path_str,
|
"db_path": db_path_str,
|
||||||
"db_exists": Path(db_path_str).exists(),
|
"db_exists": Path(db_path_str).exists(),
|
||||||
}
|
}
|
||||||
|
|
@ -100,8 +100,8 @@ if __name__ == "__main__":
|
||||||
|
|
||||||
print("Starting haiku.rag research assistant backend...")
|
print("Starting haiku.rag research assistant backend...")
|
||||||
print(f"Agent model: {agent.model}")
|
print(f"Agent model: {agent.model}")
|
||||||
print(f"QA provider: {Config.QA_PROVIDER}")
|
print(f"QA provider: {Config.qa.provider}")
|
||||||
print(f"QA model: {Config.QA_MODEL}")
|
print(f"QA model: {Config.qa.model}")
|
||||||
|
|
||||||
uvicorn.run(
|
uvicorn.run(
|
||||||
"main:app",
|
"main:app",
|
||||||
|
|
|
||||||
|
|
@ -6,16 +6,15 @@ services:
|
||||||
ports:
|
ports:
|
||||||
- "8000:8000"
|
- "8000:8000"
|
||||||
environment:
|
environment:
|
||||||
- QA_PROVIDER=${QA_PROVIDER:-ollama}
|
|
||||||
- QA_MODEL=${QA_MODEL:-gpt-oss:latest}
|
|
||||||
- DB_PATH=/app/data/haiku.rag.lancedb
|
- DB_PATH=/app/data/haiku.rag.lancedb
|
||||||
- OLLAMA_BASE_URL=${OLLAMA_BASE_URL:-http://host.docker.internal:11434}
|
# API keys (set these in your shell or .env file)
|
||||||
- OPENAI_API_KEY=${OPENAI_API_KEY}
|
- OPENAI_API_KEY=${OPENAI_API_KEY}
|
||||||
- ANTHROPIC_API_KEY=${ANTHROPIC_API_KEY}
|
- ANTHROPIC_API_KEY=${ANTHROPIC_API_KEY}
|
||||||
volumes:
|
volumes:
|
||||||
- ./backend:/app
|
- ./backend:/app
|
||||||
- /app/.venv
|
- /app/.venv
|
||||||
- ${DB_PATH}:/app/data/haiku.rag.lancedb
|
- ${DB_PATH}:/app/data/haiku.rag.lancedb
|
||||||
|
- ./haiku.rag.yaml:/app/haiku.rag.yaml:ro # Mount config file
|
||||||
networks:
|
networks:
|
||||||
- ag-ui-network
|
- ag-ui-network
|
||||||
extra_hosts:
|
extra_hosts:
|
||||||
|
|
|
||||||
20
examples/ag-ui-research/haiku.rag.yaml.example
Normal file
20
examples/ag-ui-research/haiku.rag.yaml.example
Normal file
|
|
@ -0,0 +1,20 @@
|
||||||
|
# haiku.rag configuration for ag-ui-research example
|
||||||
|
# Copy to haiku.rag.yaml and customize
|
||||||
|
|
||||||
|
qa:
|
||||||
|
provider: ollama
|
||||||
|
model: gpt-oss:latest
|
||||||
|
|
||||||
|
providers:
|
||||||
|
ollama:
|
||||||
|
base_url: http://host.docker.internal:11434
|
||||||
|
|
||||||
|
# For OpenAI:
|
||||||
|
# qa:
|
||||||
|
# provider: openai
|
||||||
|
# model: gpt-4o-mini
|
||||||
|
|
||||||
|
# For Anthropic:
|
||||||
|
# qa:
|
||||||
|
# provider: anthropic
|
||||||
|
# model: claude-3-5-haiku-20241022
|
||||||
|
|
@ -6,7 +6,7 @@ Run haiku.rag with file monitoring, MCP server, and A2A agent.
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
mkdir -p data docs
|
mkdir -p data docs
|
||||||
cp .env.example .env # Edit if needed
|
cp haiku.rag.yaml.example haiku.rag.yaml # Edit as needed
|
||||||
docker compose up -d
|
docker compose up -d
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
@ -35,10 +35,18 @@ docker compose exec haiku-rag haiku-rag a2aclient --url http://localhost:8000
|
||||||
|
|
||||||
## Configuration
|
## Configuration
|
||||||
|
|
||||||
Edit `.env` or `docker-compose.yml` to configure providers. See the [Configuration documentation](https://ggozad.github.io/haiku.rag/configuration/) for all options.
|
Edit `haiku.rag.yaml` to configure providers, embeddings, and other settings. See the [Configuration documentation](https://ggozad.github.io/haiku.rag/configuration/) for all options.
|
||||||
|
|
||||||
Default setup uses Ollama on the host (`host.docker.internal:11434`).
|
Default setup uses Ollama on the host (`host.docker.internal:11434`).
|
||||||
|
|
||||||
|
For API keys (OpenAI, Anthropic, etc.), set them as environment variables:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
export OPENAI_API_KEY=your-key-here
|
||||||
|
export ANTHROPIC_API_KEY=your-key-here
|
||||||
|
docker compose up -d
|
||||||
|
```
|
||||||
|
|
||||||
## Documentation
|
## Documentation
|
||||||
|
|
||||||
- [Configuration](https://ggozad.github.io/haiku.rag/configuration/)
|
- [Configuration](https://ggozad.github.io/haiku.rag/configuration/)
|
||||||
|
|
|
||||||
|
|
@ -10,46 +10,13 @@ services:
|
||||||
volumes:
|
volumes:
|
||||||
- ./data:/data # Persist database
|
- ./data:/data # Persist database
|
||||||
- ./docs:/docs # Mount documents directory for monitoring
|
- ./docs:/docs # Mount documents directory for monitoring
|
||||||
|
- ./haiku.rag.yaml:/app/haiku.rag.yaml:ro # Mount config file
|
||||||
environment:
|
environment:
|
||||||
# Database directory
|
# API keys (set as needed)
|
||||||
- DEFAULT_DATA_DIR=/data
|
- OPENAI_API_KEY=${OPENAI_API_KEY}
|
||||||
|
- ANTHROPIC_API_KEY=${ANTHROPIC_API_KEY}
|
||||||
# File monitoring
|
- VOYAGE_API_KEY=${VOYAGE_API_KEY}
|
||||||
- MONITOR_DIRECTORIES=/docs
|
- CO_API_KEY=${CO_API_KEY}
|
||||||
|
|
||||||
# Set the Ollama base url for Ollama defaults
|
|
||||||
- OLLAMA_BASE_URL=http://host.docker.internal:11434
|
|
||||||
|
|
||||||
# Embeddings provider (choose one)
|
|
||||||
# For OpenAI:
|
|
||||||
# - EMBEDDINGS_PROVIDER=openai
|
|
||||||
# - EMBEDDINGS_MODEL=text-embedding-3-small
|
|
||||||
# - OPENAI_API_KEY=your-key-here
|
|
||||||
|
|
||||||
# For VoyageAI:
|
|
||||||
# - EMBEDDINGS_PROVIDER=voyageai
|
|
||||||
# - EMBEDDINGS_MODEL=voyage-3
|
|
||||||
# - VOYAGE_API_KEY=your-key-here
|
|
||||||
|
|
||||||
# QA provider (uses Pydantic AI)
|
|
||||||
# - QA_PROVIDER=ollama
|
|
||||||
|
|
||||||
# For OpenAI:
|
|
||||||
# - QA_PROVIDER=openai
|
|
||||||
# - QA_MODEL=gpt-4o-mini
|
|
||||||
# - OPENAI_API_KEY=your-key-here
|
|
||||||
|
|
||||||
# Reranking (optional)
|
|
||||||
# - RERANKING_PROVIDER=mxbai
|
|
||||||
# - RERANKING_MODEL=mixedbread-ai/mxbai-rerank-large-v1
|
|
||||||
# - RERANKING_BASE_URL=http://host.docker.internal:11434
|
|
||||||
|
|
||||||
# Research agent (optional, defaults to QA provider/model)
|
|
||||||
# - RESEARCH_PROVIDER=openai
|
|
||||||
# - RESEARCH_MODEL=gpt-4o
|
|
||||||
|
|
||||||
# Environment
|
|
||||||
- ENV=production
|
|
||||||
|
|
||||||
restart: unless-stopped
|
restart: unless-stopped
|
||||||
healthcheck:
|
healthcheck:
|
||||||
|
|
|
||||||
25
examples/docker/haiku.rag.yaml.example
Normal file
25
examples/docker/haiku.rag.yaml.example
Normal file
|
|
@ -0,0 +1,25 @@
|
||||||
|
# haiku.rag configuration for Docker deployment
|
||||||
|
# See https://ggozad.github.io/haiku.rag/configuration/ for details
|
||||||
|
|
||||||
|
environment: production
|
||||||
|
|
||||||
|
storage:
|
||||||
|
data_dir: /data
|
||||||
|
monitor_directories:
|
||||||
|
- /docs
|
||||||
|
|
||||||
|
embeddings:
|
||||||
|
provider: ollama
|
||||||
|
model: nomic-embed-text
|
||||||
|
vector_dim: 768
|
||||||
|
|
||||||
|
qa:
|
||||||
|
provider: ollama
|
||||||
|
model: qwen3
|
||||||
|
|
||||||
|
providers:
|
||||||
|
ollama:
|
||||||
|
base_url: http://host.docker.internal:11434
|
||||||
|
|
||||||
|
# For other providers (OpenAI, Anthropic, VoyageAI, etc.),
|
||||||
|
# see: https://ggozad.github.io/haiku.rag/configuration/
|
||||||
|
|
@ -30,6 +30,7 @@ dependencies = [
|
||||||
"pydantic-ai>=1.0.18",
|
"pydantic-ai>=1.0.18",
|
||||||
"pydantic-graph>=1.0.18",
|
"pydantic-graph>=1.0.18",
|
||||||
"python-dotenv>=1.1.1",
|
"python-dotenv>=1.1.1",
|
||||||
|
"pyyaml>=6.0.1",
|
||||||
"rich>=14.2.0",
|
"rich>=14.2.0",
|
||||||
"tiktoken>=0.12.0",
|
"tiktoken>=0.12.0",
|
||||||
"typer>=0.19.2",
|
"typer>=0.19.2",
|
||||||
|
|
|
||||||
|
|
@ -174,7 +174,7 @@ async def run_qa_benchmark(
|
||||||
|
|
||||||
judge_model = OpenAIChatModel(
|
judge_model = OpenAIChatModel(
|
||||||
model_name=QA_JUDGE_MODEL,
|
model_name=QA_JUDGE_MODEL,
|
||||||
provider=OllamaProvider(base_url=f"{Config.OLLAMA_BASE_URL}/v1"),
|
provider=OllamaProvider(base_url=f"{Config.providers.ollama.base_url}/v1"),
|
||||||
)
|
)
|
||||||
|
|
||||||
evaluation_dataset = EvalDataset[str, str, dict[str, str]](
|
evaluation_dataset = EvalDataset[str, str, dict[str, str]](
|
||||||
|
|
|
||||||
|
|
@ -41,7 +41,7 @@ class LLMJudge:
|
||||||
# Create Ollama model
|
# Create Ollama model
|
||||||
ollama_model = OpenAIChatModel(
|
ollama_model = OpenAIChatModel(
|
||||||
model_name=model,
|
model_name=model,
|
||||||
provider=OllamaProvider(base_url=f"{Config.OLLAMA_BASE_URL}/v1"),
|
provider=OllamaProvider(base_url=f"{Config.providers.ollama.base_url}/v1"),
|
||||||
)
|
)
|
||||||
|
|
||||||
# Create Pydantic AI agent
|
# Create Pydantic AI agent
|
||||||
|
|
|
||||||
|
|
@ -57,12 +57,12 @@ def create_a2a_app(
|
||||||
"""
|
"""
|
||||||
base_storage = InMemoryStorage()
|
base_storage = InMemoryStorage()
|
||||||
storage = LRUMemoryStorage(
|
storage = LRUMemoryStorage(
|
||||||
storage=base_storage, max_contexts=Config.A2A_MAX_CONTEXTS
|
storage=base_storage, max_contexts=Config.a2a.max_contexts
|
||||||
)
|
)
|
||||||
broker = InMemoryBroker()
|
broker = InMemoryBroker()
|
||||||
|
|
||||||
# Create the agent with native search tool
|
# Create the agent with native search tool
|
||||||
model = get_model(Config.QA_PROVIDER, Config.QA_MODEL)
|
model = get_model(Config.qa.provider, Config.qa.model)
|
||||||
agent = Agent(
|
agent = Agent(
|
||||||
model=model,
|
model=model,
|
||||||
deps_type=AgentDependencies,
|
deps_type=AgentDependencies,
|
||||||
|
|
@ -120,7 +120,7 @@ def create_a2a_app(
|
||||||
# Create FastA2A app with custom worker lifecycle
|
# Create FastA2A app with custom worker lifecycle
|
||||||
@asynccontextmanager
|
@asynccontextmanager
|
||||||
async def lifespan(app):
|
async def lifespan(app):
|
||||||
logger.info(f"Started A2A server (max contexts: {Config.A2A_MAX_CONTEXTS})")
|
logger.info(f"Started A2A server (max contexts: {Config.a2a.max_contexts})")
|
||||||
async with app.task_manager:
|
async with app.task_manager:
|
||||||
async with worker.run():
|
async with worker.run():
|
||||||
yield
|
yield
|
||||||
|
|
|
||||||
|
|
@ -231,8 +231,8 @@ class HaikuRAGApp:
|
||||||
)
|
)
|
||||||
|
|
||||||
start_node = DeepQAPlanNode(
|
start_node = DeepQAPlanNode(
|
||||||
provider=Config.QA_PROVIDER,
|
provider=Config.qa.provider,
|
||||||
model=Config.QA_MODEL,
|
model=Config.qa.model,
|
||||||
)
|
)
|
||||||
|
|
||||||
result = await graph.run(
|
result = await graph.run(
|
||||||
|
|
@ -278,8 +278,8 @@ class HaikuRAGApp:
|
||||||
)
|
)
|
||||||
|
|
||||||
start = PlanNode(
|
start = PlanNode(
|
||||||
provider=Config.RESEARCH_PROVIDER or Config.QA_PROVIDER,
|
provider=Config.research.provider or Config.qa.provider,
|
||||||
model=Config.RESEARCH_MODEL or Config.QA_MODEL,
|
model=Config.research.model or Config.qa.model,
|
||||||
)
|
)
|
||||||
report = None
|
report = None
|
||||||
async for event in stream_research_graph(graph, start, state, deps):
|
async for event in stream_research_graph(graph, start, state, deps):
|
||||||
|
|
@ -474,7 +474,9 @@ class HaikuRAGApp:
|
||||||
|
|
||||||
# Start file monitor if enabled
|
# Start file monitor if enabled
|
||||||
if enable_monitor:
|
if enable_monitor:
|
||||||
monitor = FileWatcher(paths=Config.MONITOR_DIRECTORIES, client=client)
|
monitor = FileWatcher(
|
||||||
|
paths=Config.storage.monitor_directories, client=client
|
||||||
|
)
|
||||||
monitor_task = asyncio.create_task(monitor.observe())
|
monitor_task = asyncio.create_task(monitor.observe())
|
||||||
tasks.append(monitor_task)
|
tasks.append(monitor_task)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -22,7 +22,7 @@ class Chunker:
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
chunk_size: int = Config.CHUNK_SIZE,
|
chunk_size: int = Config.processing.chunk_size,
|
||||||
):
|
):
|
||||||
self.chunk_size = chunk_size
|
self.chunk_size = chunk_size
|
||||||
tokenizer = OpenAITokenizer(
|
tokenizer = OpenAITokenizer(
|
||||||
|
|
|
||||||
|
|
@ -42,10 +42,21 @@ def main(
|
||||||
callback=version_callback,
|
callback=version_callback,
|
||||||
help="Show version and exit",
|
help="Show version and exit",
|
||||||
),
|
),
|
||||||
|
config: Path | None = typer.Option(
|
||||||
|
None,
|
||||||
|
"--config",
|
||||||
|
help="Path to YAML configuration file",
|
||||||
|
),
|
||||||
):
|
):
|
||||||
"""haiku.rag CLI - Vector database RAG system"""
|
"""haiku.rag CLI - Vector database RAG system"""
|
||||||
|
# Store config path in environment for config loader to use
|
||||||
|
if config:
|
||||||
|
import os
|
||||||
|
|
||||||
|
os.environ["HAIKU_RAG_CONFIG_PATH"] = str(config.absolute())
|
||||||
|
|
||||||
# Configure logging minimally for CLI context
|
# Configure logging minimally for CLI context
|
||||||
if Config.ENV == "development":
|
if Config.environment == "development":
|
||||||
# Lazy import logfire only in development
|
# Lazy import logfire only in development
|
||||||
try:
|
try:
|
||||||
import logfire # type: ignore
|
import logfire # type: ignore
|
||||||
|
|
@ -69,7 +80,7 @@ def main(
|
||||||
@cli.command("list", help="List all stored documents")
|
@cli.command("list", help="List all stored documents")
|
||||||
def list_documents(
|
def list_documents(
|
||||||
db: Path = typer.Option(
|
db: Path = typer.Option(
|
||||||
Config.DEFAULT_DATA_DIR / "haiku.rag.lancedb",
|
Config.storage.data_dir / "haiku.rag.lancedb",
|
||||||
"--db",
|
"--db",
|
||||||
help="Path to the LanceDB database file",
|
help="Path to the LanceDB database file",
|
||||||
),
|
),
|
||||||
|
|
@ -116,7 +127,7 @@ def add_document_text(
|
||||||
metavar="KEY=VALUE",
|
metavar="KEY=VALUE",
|
||||||
),
|
),
|
||||||
db: Path = typer.Option(
|
db: Path = typer.Option(
|
||||||
Config.DEFAULT_DATA_DIR / "haiku.rag.lancedb",
|
Config.storage.data_dir / "haiku.rag.lancedb",
|
||||||
"--db",
|
"--db",
|
||||||
help="Path to the LanceDB database file",
|
help="Path to the LanceDB database file",
|
||||||
),
|
),
|
||||||
|
|
@ -145,7 +156,7 @@ def add_document_src(
|
||||||
metavar="KEY=VALUE",
|
metavar="KEY=VALUE",
|
||||||
),
|
),
|
||||||
db: Path = typer.Option(
|
db: Path = typer.Option(
|
||||||
Config.DEFAULT_DATA_DIR / "haiku.rag.lancedb",
|
Config.storage.data_dir / "haiku.rag.lancedb",
|
||||||
"--db",
|
"--db",
|
||||||
help="Path to the LanceDB database file",
|
help="Path to the LanceDB database file",
|
||||||
),
|
),
|
||||||
|
|
@ -167,7 +178,7 @@ def get_document(
|
||||||
help="The ID of the document to get",
|
help="The ID of the document to get",
|
||||||
),
|
),
|
||||||
db: Path = typer.Option(
|
db: Path = typer.Option(
|
||||||
Config.DEFAULT_DATA_DIR / "haiku.rag.lancedb",
|
Config.storage.data_dir / "haiku.rag.lancedb",
|
||||||
"--db",
|
"--db",
|
||||||
help="Path to the LanceDB database file",
|
help="Path to the LanceDB database file",
|
||||||
),
|
),
|
||||||
|
|
@ -184,7 +195,7 @@ def delete_document(
|
||||||
help="The ID of the document to delete",
|
help="The ID of the document to delete",
|
||||||
),
|
),
|
||||||
db: Path = typer.Option(
|
db: Path = typer.Option(
|
||||||
Config.DEFAULT_DATA_DIR / "haiku.rag.lancedb",
|
Config.storage.data_dir / "haiku.rag.lancedb",
|
||||||
"--db",
|
"--db",
|
||||||
help="Path to the LanceDB database file",
|
help="Path to the LanceDB database file",
|
||||||
),
|
),
|
||||||
|
|
@ -211,7 +222,7 @@ def search(
|
||||||
help="Maximum number of results to return",
|
help="Maximum number of results to return",
|
||||||
),
|
),
|
||||||
db: Path = typer.Option(
|
db: Path = typer.Option(
|
||||||
Config.DEFAULT_DATA_DIR / "haiku.rag.lancedb",
|
Config.storage.data_dir / "haiku.rag.lancedb",
|
||||||
"--db",
|
"--db",
|
||||||
help="Path to the LanceDB database file",
|
help="Path to the LanceDB database file",
|
||||||
),
|
),
|
||||||
|
|
@ -228,7 +239,7 @@ def ask(
|
||||||
help="The question to ask",
|
help="The question to ask",
|
||||||
),
|
),
|
||||||
db: Path = typer.Option(
|
db: Path = typer.Option(
|
||||||
Config.DEFAULT_DATA_DIR / "haiku.rag.lancedb",
|
Config.storage.data_dir / "haiku.rag.lancedb",
|
||||||
"--db",
|
"--db",
|
||||||
help="Path to the LanceDB database file",
|
help="Path to the LanceDB database file",
|
||||||
),
|
),
|
||||||
|
|
@ -276,7 +287,7 @@ def research(
|
||||||
help="Max concurrent searches per iteration (planned)",
|
help="Max concurrent searches per iteration (planned)",
|
||||||
),
|
),
|
||||||
db: Path = typer.Option(
|
db: Path = typer.Option(
|
||||||
Config.DEFAULT_DATA_DIR / "haiku.rag.lancedb",
|
Config.storage.data_dir / "haiku.rag.lancedb",
|
||||||
"--db",
|
"--db",
|
||||||
help="Path to the LanceDB database file",
|
help="Path to the LanceDB database file",
|
||||||
),
|
),
|
||||||
|
|
@ -308,13 +319,61 @@ def settings():
|
||||||
app.show_settings()
|
app.show_settings()
|
||||||
|
|
||||||
|
|
||||||
|
@cli.command("init-config", help="Generate a YAML configuration file")
|
||||||
|
def init_config(
|
||||||
|
output: Path = typer.Argument(
|
||||||
|
Path("haiku.rag.yaml"),
|
||||||
|
help="Output path for the config file",
|
||||||
|
),
|
||||||
|
from_env: bool = typer.Option(
|
||||||
|
False,
|
||||||
|
"--from-env",
|
||||||
|
help="Migrate settings from .env file",
|
||||||
|
),
|
||||||
|
):
|
||||||
|
"""Generate a YAML configuration file with defaults or from .env."""
|
||||||
|
import yaml
|
||||||
|
|
||||||
|
from haiku.rag.config.loader import generate_default_config, load_config_from_env
|
||||||
|
|
||||||
|
if output.exists():
|
||||||
|
typer.echo(
|
||||||
|
f"Error: {output} already exists. Remove it first or choose a different path."
|
||||||
|
)
|
||||||
|
raise typer.Exit(1)
|
||||||
|
|
||||||
|
if from_env:
|
||||||
|
# Load from environment variables (including .env if present)
|
||||||
|
from dotenv import load_dotenv
|
||||||
|
|
||||||
|
load_dotenv()
|
||||||
|
config_data = load_config_from_env()
|
||||||
|
if not config_data:
|
||||||
|
typer.echo("Warning: No environment variables found to migrate.")
|
||||||
|
typer.echo("Generating default configuration instead.")
|
||||||
|
config_data = generate_default_config()
|
||||||
|
else:
|
||||||
|
config_data = generate_default_config()
|
||||||
|
|
||||||
|
# Write YAML with comments
|
||||||
|
with open(output, "w") as f:
|
||||||
|
f.write("# haiku.rag configuration file\n")
|
||||||
|
f.write(
|
||||||
|
"# See https://ggozad.github.io/haiku.rag/configuration/ for details\n\n"
|
||||||
|
)
|
||||||
|
yaml.dump(config_data, f, default_flow_style=False, sort_keys=False)
|
||||||
|
|
||||||
|
typer.echo(f"Configuration file created: {output}")
|
||||||
|
typer.echo("Edit the file to customize your settings.")
|
||||||
|
|
||||||
|
|
||||||
@cli.command(
|
@cli.command(
|
||||||
"rebuild",
|
"rebuild",
|
||||||
help="Rebuild the database by deleting all chunks and re-indexing all documents",
|
help="Rebuild the database by deleting all chunks and re-indexing all documents",
|
||||||
)
|
)
|
||||||
def rebuild(
|
def rebuild(
|
||||||
db: Path = typer.Option(
|
db: Path = typer.Option(
|
||||||
Config.DEFAULT_DATA_DIR / "haiku.rag.lancedb",
|
Config.storage.data_dir / "haiku.rag.lancedb",
|
||||||
"--db",
|
"--db",
|
||||||
help="Path to the LanceDB database file",
|
help="Path to the LanceDB database file",
|
||||||
),
|
),
|
||||||
|
|
@ -328,7 +387,7 @@ def rebuild(
|
||||||
@cli.command("vacuum", help="Optimize and clean up all tables to reduce disk usage")
|
@cli.command("vacuum", help="Optimize and clean up all tables to reduce disk usage")
|
||||||
def vacuum(
|
def vacuum(
|
||||||
db: Path = typer.Option(
|
db: Path = typer.Option(
|
||||||
Config.DEFAULT_DATA_DIR / "haiku.rag.lancedb",
|
Config.storage.data_dir / "haiku.rag.lancedb",
|
||||||
"--db",
|
"--db",
|
||||||
help="Path to the LanceDB database file",
|
help="Path to the LanceDB database file",
|
||||||
),
|
),
|
||||||
|
|
@ -342,7 +401,7 @@ def vacuum(
|
||||||
@cli.command("info", help="Show read-only database info (no upgrades or writes)")
|
@cli.command("info", help="Show read-only database info (no upgrades or writes)")
|
||||||
def info(
|
def info(
|
||||||
db: Path = typer.Option(
|
db: Path = typer.Option(
|
||||||
Config.DEFAULT_DATA_DIR / "haiku.rag.lancedb",
|
Config.storage.data_dir / "haiku.rag.lancedb",
|
||||||
"--db",
|
"--db",
|
||||||
help="Path to the LanceDB database file",
|
help="Path to the LanceDB database file",
|
||||||
),
|
),
|
||||||
|
|
@ -371,7 +430,7 @@ def download_models_cmd():
|
||||||
)
|
)
|
||||||
def serve(
|
def serve(
|
||||||
db: Path = typer.Option(
|
db: Path = typer.Option(
|
||||||
Config.DEFAULT_DATA_DIR / "haiku.rag.lancedb",
|
Config.storage.data_dir / "haiku.rag.lancedb",
|
||||||
"--db",
|
"--db",
|
||||||
help="Path to the LanceDB database file",
|
help="Path to the LanceDB database 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.DEFAULT_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.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]:
|
||||||
|
|
|
||||||
|
|
@ -1,90 +0,0 @@
|
||||||
import os
|
|
||||||
from pathlib import Path
|
|
||||||
|
|
||||||
from dotenv import load_dotenv
|
|
||||||
from pydantic import BaseModel, field_validator
|
|
||||||
|
|
||||||
from haiku.rag.utils import get_default_data_dir
|
|
||||||
|
|
||||||
load_dotenv()
|
|
||||||
|
|
||||||
|
|
||||||
class AppConfig(BaseModel):
|
|
||||||
ENV: str = "production"
|
|
||||||
|
|
||||||
LANCEDB_API_KEY: str = ""
|
|
||||||
LANCEDB_URI: str = ""
|
|
||||||
LANCEDB_REGION: str = ""
|
|
||||||
|
|
||||||
DEFAULT_DATA_DIR: Path = get_default_data_dir()
|
|
||||||
MONITOR_DIRECTORIES: list[Path] = []
|
|
||||||
|
|
||||||
EMBEDDINGS_PROVIDER: str = "ollama"
|
|
||||||
EMBEDDINGS_MODEL: str = "qwen3-embedding"
|
|
||||||
EMBEDDINGS_VECTOR_DIM: int = 4096
|
|
||||||
|
|
||||||
RERANK_PROVIDER: str = ""
|
|
||||||
RERANK_MODEL: str = ""
|
|
||||||
|
|
||||||
QA_PROVIDER: str = "ollama"
|
|
||||||
QA_MODEL: str = "gpt-oss"
|
|
||||||
|
|
||||||
# Research defaults (fallback to QA if not provided via env)
|
|
||||||
RESEARCH_PROVIDER: str = "ollama"
|
|
||||||
RESEARCH_MODEL: str = "gpt-oss"
|
|
||||||
|
|
||||||
CHUNK_SIZE: int = 256
|
|
||||||
CONTEXT_CHUNK_RADIUS: int = 0
|
|
||||||
|
|
||||||
# Optional dotted path or file path to a callable that preprocesses
|
|
||||||
# markdown content before chunking. Examples:
|
|
||||||
MARKDOWN_PREPROCESSOR: str = ""
|
|
||||||
|
|
||||||
OLLAMA_BASE_URL: str = "http://localhost:11434"
|
|
||||||
|
|
||||||
VLLM_EMBEDDINGS_BASE_URL: str = ""
|
|
||||||
VLLM_RERANK_BASE_URL: str = ""
|
|
||||||
VLLM_QA_BASE_URL: str = ""
|
|
||||||
VLLM_RESEARCH_BASE_URL: str = ""
|
|
||||||
|
|
||||||
# Provider keys
|
|
||||||
VOYAGE_API_KEY: str = ""
|
|
||||||
OPENAI_API_KEY: str = ""
|
|
||||||
ANTHROPIC_API_KEY: str = ""
|
|
||||||
COHERE_API_KEY: str = ""
|
|
||||||
|
|
||||||
# If true, refuse to auto-create a new LanceDB database or tables
|
|
||||||
# and error out when the database does not already exist.
|
|
||||||
DISABLE_DB_AUTOCREATE: bool = False
|
|
||||||
|
|
||||||
# Vacuum retention threshold in seconds. Only versions older than this
|
|
||||||
# threshold will be removed during vacuum operations. Default is 60 seconds
|
|
||||||
# to allow concurrent connections to safely use recent versions.
|
|
||||||
VACUUM_RETENTION_SECONDS: int = 60
|
|
||||||
|
|
||||||
# Maximum number of A2A contexts to keep in memory. When exceeded, least
|
|
||||||
# recently used contexts will be evicted. Default is 1000.
|
|
||||||
A2A_MAX_CONTEXTS: int = 1000
|
|
||||||
|
|
||||||
@field_validator("MONITOR_DIRECTORIES", mode="before")
|
|
||||||
@classmethod
|
|
||||||
def parse_monitor_directories(cls, v):
|
|
||||||
if isinstance(v, str):
|
|
||||||
if not v.strip():
|
|
||||||
return []
|
|
||||||
return [
|
|
||||||
Path(path.strip()).absolute() for path in v.split(",") if path.strip()
|
|
||||||
]
|
|
||||||
return v
|
|
||||||
|
|
||||||
|
|
||||||
# Expose Config object for app to import
|
|
||||||
Config = AppConfig.model_validate(os.environ)
|
|
||||||
if Config.OPENAI_API_KEY:
|
|
||||||
os.environ["OPENAI_API_KEY"] = Config.OPENAI_API_KEY
|
|
||||||
if Config.VOYAGE_API_KEY:
|
|
||||||
os.environ["VOYAGE_API_KEY"] = Config.VOYAGE_API_KEY
|
|
||||||
if Config.ANTHROPIC_API_KEY:
|
|
||||||
os.environ["ANTHROPIC_API_KEY"] = Config.ANTHROPIC_API_KEY
|
|
||||||
if Config.COHERE_API_KEY:
|
|
||||||
os.environ["CO_API_KEY"] = Config.COHERE_API_KEY
|
|
||||||
54
src/haiku/rag/config/__init__.py
Normal file
54
src/haiku/rag/config/__init__.py
Normal file
|
|
@ -0,0 +1,54 @@
|
||||||
|
import os
|
||||||
|
|
||||||
|
from haiku.rag.config.loader import (
|
||||||
|
check_for_deprecated_env,
|
||||||
|
find_config_file,
|
||||||
|
generate_default_config,
|
||||||
|
load_config_from_env,
|
||||||
|
load_yaml_config,
|
||||||
|
)
|
||||||
|
from haiku.rag.config.models import (
|
||||||
|
A2AConfig,
|
||||||
|
AppConfig,
|
||||||
|
EmbeddingsConfig,
|
||||||
|
LanceDBConfig,
|
||||||
|
OllamaConfig,
|
||||||
|
ProcessingConfig,
|
||||||
|
ProvidersConfig,
|
||||||
|
QAConfig,
|
||||||
|
RerankingConfig,
|
||||||
|
ResearchConfig,
|
||||||
|
StorageConfig,
|
||||||
|
VLLMConfig,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"Config",
|
||||||
|
"AppConfig",
|
||||||
|
"StorageConfig",
|
||||||
|
"LanceDBConfig",
|
||||||
|
"EmbeddingsConfig",
|
||||||
|
"RerankingConfig",
|
||||||
|
"QAConfig",
|
||||||
|
"ResearchConfig",
|
||||||
|
"ProcessingConfig",
|
||||||
|
"OllamaConfig",
|
||||||
|
"VLLMConfig",
|
||||||
|
"ProvidersConfig",
|
||||||
|
"A2AConfig",
|
||||||
|
"find_config_file",
|
||||||
|
"load_yaml_config",
|
||||||
|
"generate_default_config",
|
||||||
|
"load_config_from_env",
|
||||||
|
]
|
||||||
|
|
||||||
|
# Load config from YAML file or use defaults
|
||||||
|
config_path = find_config_file(None)
|
||||||
|
if config_path:
|
||||||
|
yaml_data = load_yaml_config(config_path)
|
||||||
|
Config = AppConfig.model_validate(yaml_data)
|
||||||
|
else:
|
||||||
|
Config = AppConfig()
|
||||||
|
|
||||||
|
# Check for deprecated .env file
|
||||||
|
check_for_deprecated_env()
|
||||||
151
src/haiku/rag/config/loader.py
Normal file
151
src/haiku/rag/config/loader.py
Normal file
|
|
@ -0,0 +1,151 @@
|
||||||
|
import os
|
||||||
|
import warnings
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import yaml
|
||||||
|
|
||||||
|
|
||||||
|
def find_config_file(cli_path: Path | None = None) -> Path | None:
|
||||||
|
"""Find the YAML config file using the search path.
|
||||||
|
|
||||||
|
Search order:
|
||||||
|
1. CLI-provided path (via HAIKU_RAG_CONFIG_PATH env var or parameter)
|
||||||
|
2. ./haiku.rag.yaml (current directory)
|
||||||
|
3. ~/.config/haiku.rag/config.yaml (user config)
|
||||||
|
|
||||||
|
Returns None if no config file is found.
|
||||||
|
"""
|
||||||
|
# Check environment variable first (set by CLI --config flag)
|
||||||
|
if not cli_path:
|
||||||
|
env_path = os.getenv("HAIKU_RAG_CONFIG_PATH")
|
||||||
|
if env_path:
|
||||||
|
cli_path = Path(env_path)
|
||||||
|
|
||||||
|
if cli_path:
|
||||||
|
if cli_path.exists():
|
||||||
|
return cli_path
|
||||||
|
raise FileNotFoundError(f"Config file not found: {cli_path}")
|
||||||
|
|
||||||
|
cwd_config = Path.cwd() / "haiku.rag.yaml"
|
||||||
|
if cwd_config.exists():
|
||||||
|
return cwd_config
|
||||||
|
|
||||||
|
user_config_dir = Path.home() / ".config" / "haiku.rag"
|
||||||
|
user_config = user_config_dir / "config.yaml"
|
||||||
|
if user_config.exists():
|
||||||
|
return user_config
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def load_yaml_config(path: Path) -> dict:
|
||||||
|
"""Load and parse a YAML config file."""
|
||||||
|
with open(path) as f:
|
||||||
|
data = yaml.safe_load(f)
|
||||||
|
return data or {}
|
||||||
|
|
||||||
|
|
||||||
|
def check_for_deprecated_env() -> None:
|
||||||
|
"""Check for .env file and warn if found."""
|
||||||
|
env_file = Path.cwd() / ".env"
|
||||||
|
if env_file.exists():
|
||||||
|
warnings.warn(
|
||||||
|
".env file detected but YAML configuration is now preferred. "
|
||||||
|
"Environment variable configuration is deprecated and will be removed in future versions."
|
||||||
|
"Run 'haiku-rag init-config' to generate a YAML config file.",
|
||||||
|
DeprecationWarning,
|
||||||
|
stacklevel=2,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def generate_default_config() -> dict:
|
||||||
|
"""Generate a default YAML config structure with documentation."""
|
||||||
|
return {
|
||||||
|
"environment": "production",
|
||||||
|
"storage": {
|
||||||
|
"data_dir": "",
|
||||||
|
"monitor_directories": [],
|
||||||
|
"disable_autocreate": False,
|
||||||
|
"vacuum_retention_seconds": 60,
|
||||||
|
},
|
||||||
|
"lancedb": {"uri": "", "api_key": "", "region": ""},
|
||||||
|
"embeddings": {
|
||||||
|
"provider": "ollama",
|
||||||
|
"model": "qwen3-embedding",
|
||||||
|
"vector_dim": 4096,
|
||||||
|
},
|
||||||
|
"reranking": {"provider": "", "model": ""},
|
||||||
|
"qa": {"provider": "ollama", "model": "gpt-oss"},
|
||||||
|
"research": {"provider": "", "model": ""},
|
||||||
|
"processing": {
|
||||||
|
"chunk_size": 256,
|
||||||
|
"context_chunk_radius": 0,
|
||||||
|
"markdown_preprocessor": "",
|
||||||
|
},
|
||||||
|
"providers": {
|
||||||
|
"ollama": {"base_url": "http://localhost:11434"},
|
||||||
|
"vllm": {
|
||||||
|
"embeddings_base_url": "",
|
||||||
|
"rerank_base_url": "",
|
||||||
|
"qa_base_url": "",
|
||||||
|
"research_base_url": "",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"a2a": {"max_contexts": 1000},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def load_config_from_env() -> dict:
|
||||||
|
"""Load current config from environment variables (for migration)."""
|
||||||
|
result = {}
|
||||||
|
|
||||||
|
env_mappings = {
|
||||||
|
"ENV": "environment",
|
||||||
|
"DEFAULT_DATA_DIR": ("storage", "data_dir"),
|
||||||
|
"MONITOR_DIRECTORIES": ("storage", "monitor_directories"),
|
||||||
|
"DISABLE_DB_AUTOCREATE": ("storage", "disable_autocreate"),
|
||||||
|
"VACUUM_RETENTION_SECONDS": ("storage", "vacuum_retention_seconds"),
|
||||||
|
"LANCEDB_URI": ("lancedb", "uri"),
|
||||||
|
"LANCEDB_API_KEY": ("lancedb", "api_key"),
|
||||||
|
"LANCEDB_REGION": ("lancedb", "region"),
|
||||||
|
"EMBEDDINGS_PROVIDER": ("embeddings", "provider"),
|
||||||
|
"EMBEDDINGS_MODEL": ("embeddings", "model"),
|
||||||
|
"EMBEDDINGS_VECTOR_DIM": ("embeddings", "vector_dim"),
|
||||||
|
"RERANK_PROVIDER": ("reranking", "provider"),
|
||||||
|
"RERANK_MODEL": ("reranking", "model"),
|
||||||
|
"QA_PROVIDER": ("qa", "provider"),
|
||||||
|
"QA_MODEL": ("qa", "model"),
|
||||||
|
"RESEARCH_PROVIDER": ("research", "provider"),
|
||||||
|
"RESEARCH_MODEL": ("research", "model"),
|
||||||
|
"CHUNK_SIZE": ("processing", "chunk_size"),
|
||||||
|
"CONTEXT_CHUNK_RADIUS": ("processing", "context_chunk_radius"),
|
||||||
|
"MARKDOWN_PREPROCESSOR": ("processing", "markdown_preprocessor"),
|
||||||
|
"OLLAMA_BASE_URL": ("providers", "ollama", "base_url"),
|
||||||
|
"VLLM_EMBEDDINGS_BASE_URL": ("providers", "vllm", "embeddings_base_url"),
|
||||||
|
"VLLM_RERANK_BASE_URL": ("providers", "vllm", "rerank_base_url"),
|
||||||
|
"VLLM_QA_BASE_URL": ("providers", "vllm", "qa_base_url"),
|
||||||
|
"VLLM_RESEARCH_BASE_URL": ("providers", "vllm", "research_base_url"),
|
||||||
|
"A2A_MAX_CONTEXTS": ("a2a", "max_contexts"),
|
||||||
|
}
|
||||||
|
|
||||||
|
for env_var, path in env_mappings.items():
|
||||||
|
value = os.getenv(env_var)
|
||||||
|
if value is not None:
|
||||||
|
# Special handling for MONITOR_DIRECTORIES - parse comma-separated list
|
||||||
|
if env_var == "MONITOR_DIRECTORIES":
|
||||||
|
if value.strip():
|
||||||
|
value = [p.strip() for p in value.split(",") if p.strip()]
|
||||||
|
else:
|
||||||
|
value = []
|
||||||
|
|
||||||
|
if isinstance(path, tuple):
|
||||||
|
current = result
|
||||||
|
for key in path[:-1]:
|
||||||
|
if key not in current:
|
||||||
|
current[key] = {}
|
||||||
|
current = current[key]
|
||||||
|
current[path[-1]] = value
|
||||||
|
else:
|
||||||
|
result[path] = value
|
||||||
|
|
||||||
|
return result
|
||||||
78
src/haiku/rag/config/models.py
Normal file
78
src/haiku/rag/config/models.py
Normal file
|
|
@ -0,0 +1,78 @@
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
from haiku.rag.utils import get_default_data_dir
|
||||||
|
|
||||||
|
|
||||||
|
class StorageConfig(BaseModel):
|
||||||
|
data_dir: Path = Field(default_factory=get_default_data_dir)
|
||||||
|
monitor_directories: list[Path] = []
|
||||||
|
disable_autocreate: bool = False
|
||||||
|
vacuum_retention_seconds: int = 60
|
||||||
|
|
||||||
|
|
||||||
|
class LanceDBConfig(BaseModel):
|
||||||
|
uri: str = ""
|
||||||
|
api_key: str = ""
|
||||||
|
region: str = ""
|
||||||
|
|
||||||
|
|
||||||
|
class EmbeddingsConfig(BaseModel):
|
||||||
|
provider: str = "ollama"
|
||||||
|
model: str = "qwen3-embedding"
|
||||||
|
vector_dim: int = 4096
|
||||||
|
|
||||||
|
|
||||||
|
class RerankingConfig(BaseModel):
|
||||||
|
provider: str = ""
|
||||||
|
model: str = ""
|
||||||
|
|
||||||
|
|
||||||
|
class QAConfig(BaseModel):
|
||||||
|
provider: str = "ollama"
|
||||||
|
model: str = "gpt-oss"
|
||||||
|
|
||||||
|
|
||||||
|
class ResearchConfig(BaseModel):
|
||||||
|
provider: str = "ollama"
|
||||||
|
model: str = "gpt-oss"
|
||||||
|
|
||||||
|
|
||||||
|
class ProcessingConfig(BaseModel):
|
||||||
|
chunk_size: int = 256
|
||||||
|
context_chunk_radius: int = 0
|
||||||
|
markdown_preprocessor: str = ""
|
||||||
|
|
||||||
|
|
||||||
|
class OllamaConfig(BaseModel):
|
||||||
|
base_url: str = "http://localhost:11434"
|
||||||
|
|
||||||
|
|
||||||
|
class VLLMConfig(BaseModel):
|
||||||
|
embeddings_base_url: str = ""
|
||||||
|
rerank_base_url: str = ""
|
||||||
|
qa_base_url: str = ""
|
||||||
|
research_base_url: str = ""
|
||||||
|
|
||||||
|
|
||||||
|
class ProvidersConfig(BaseModel):
|
||||||
|
ollama: OllamaConfig = Field(default_factory=OllamaConfig)
|
||||||
|
vllm: VLLMConfig = Field(default_factory=VLLMConfig)
|
||||||
|
|
||||||
|
|
||||||
|
class A2AConfig(BaseModel):
|
||||||
|
max_contexts: int = 1000
|
||||||
|
|
||||||
|
|
||||||
|
class AppConfig(BaseModel):
|
||||||
|
environment: str = "production"
|
||||||
|
storage: StorageConfig = Field(default_factory=StorageConfig)
|
||||||
|
lancedb: LanceDBConfig = Field(default_factory=LanceDBConfig)
|
||||||
|
embeddings: EmbeddingsConfig = Field(default_factory=EmbeddingsConfig)
|
||||||
|
reranking: RerankingConfig = Field(default_factory=RerankingConfig)
|
||||||
|
qa: QAConfig = Field(default_factory=QAConfig)
|
||||||
|
research: ResearchConfig = Field(default_factory=ResearchConfig)
|
||||||
|
processing: ProcessingConfig = Field(default_factory=ProcessingConfig)
|
||||||
|
providers: ProvidersConfig = Field(default_factory=ProvidersConfig)
|
||||||
|
a2a: A2AConfig = Field(default_factory=A2AConfig)
|
||||||
|
|
@ -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}")
|
||||||
|
|
|
||||||
|
|
@ -4,8 +4,8 @@ from haiku.rag.config import Config
|
||||||
|
|
||||||
|
|
||||||
class EmbedderBase:
|
class EmbedderBase:
|
||||||
_model: str = Config.EMBEDDINGS_MODEL
|
_model: str = Config.embeddings.model
|
||||||
_vector_dim: int = Config.EMBEDDINGS_VECTOR_DIM
|
_vector_dim: int = Config.embeddings.vector_dim
|
||||||
|
|
||||||
def __init__(self, model: str, vector_dim: int):
|
def __init__(self, model: str, vector_dim: int):
|
||||||
self._model = model
|
self._model = model
|
||||||
|
|
|
||||||
|
|
@ -14,7 +14,9 @@ class Embedder(EmbedderBase):
|
||||||
async def embed(self, text: list[str]) -> list[list[float]]: ...
|
async def embed(self, text: list[str]) -> list[list[float]]: ...
|
||||||
|
|
||||||
async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]:
|
async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]:
|
||||||
client = AsyncOpenAI(base_url=f"{Config.OLLAMA_BASE_URL}/v1", api_key="dummy")
|
client = AsyncOpenAI(
|
||||||
|
base_url=f"{Config.providers.ollama.base_url}/v1", api_key="dummy"
|
||||||
|
)
|
||||||
if not text:
|
if not text:
|
||||||
return []
|
return []
|
||||||
response = await client.embeddings.create(
|
response = await client.embeddings.create(
|
||||||
|
|
|
||||||
|
|
@ -15,7 +15,7 @@ class Embedder(EmbedderBase):
|
||||||
|
|
||||||
async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]:
|
async def embed(self, text: str | list[str]) -> list[float] | list[list[float]]:
|
||||||
client = AsyncOpenAI(
|
client = AsyncOpenAI(
|
||||||
base_url=f"{Config.VLLM_EMBEDDINGS_BASE_URL}/v1", api_key="dummy"
|
base_url=f"{Config.providers.vllm.embeddings_base_url}/v1", api_key="dummy"
|
||||||
)
|
)
|
||||||
if not text:
|
if not text:
|
||||||
return []
|
return []
|
||||||
|
|
|
||||||
|
|
@ -15,13 +15,13 @@ def get_model(provider: str, model: str) -> Any:
|
||||||
if provider == "ollama":
|
if provider == "ollama":
|
||||||
return OpenAIChatModel(
|
return OpenAIChatModel(
|
||||||
model_name=model,
|
model_name=model,
|
||||||
provider=OllamaProvider(base_url=f"{Config.OLLAMA_BASE_URL}/v1"),
|
provider=OllamaProvider(base_url=f"{Config.providers.ollama.base_url}/v1"),
|
||||||
)
|
)
|
||||||
elif provider == "vllm":
|
elif provider == "vllm":
|
||||||
return OpenAIChatModel(
|
return OpenAIChatModel(
|
||||||
model_name=model,
|
model_name=model,
|
||||||
provider=OpenAIProvider(
|
provider=OpenAIProvider(
|
||||||
base_url=f"{Config.VLLM_RESEARCH_BASE_URL or Config.VLLM_QA_BASE_URL}/v1",
|
base_url=f"{Config.providers.vllm.research_base_url or Config.providers.vllm.qa_base_url}/v1",
|
||||||
api_key="none",
|
api_key="none",
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -38,10 +38,13 @@ def create_mcp_server(db_path: Path) -> FastMCP:
|
||||||
"""Add a document to the RAG system from a file path."""
|
"""Add a document to the RAG system from a file path."""
|
||||||
try:
|
try:
|
||||||
async with HaikuRAG(db_path) as rag:
|
async with HaikuRAG(db_path) as rag:
|
||||||
document = await rag.create_document_from_source(
|
result = await rag.create_document_from_source(
|
||||||
Path(file_path), title=title, metadata=metadata or {}
|
Path(file_path), title=title, metadata=metadata or {}
|
||||||
)
|
)
|
||||||
return document.id
|
# Handle both single document and list of documents (directories)
|
||||||
|
if isinstance(result, list):
|
||||||
|
return result[0].id if result else None
|
||||||
|
return result.id
|
||||||
except Exception:
|
except Exception:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
@ -52,10 +55,13 @@ def create_mcp_server(db_path: Path) -> FastMCP:
|
||||||
"""Add a document to the RAG system from a URL."""
|
"""Add a document to the RAG system from a URL."""
|
||||||
try:
|
try:
|
||||||
async with HaikuRAG(db_path) as rag:
|
async with HaikuRAG(db_path) as rag:
|
||||||
document = await rag.create_document_from_source(
|
result = await rag.create_document_from_source(
|
||||||
url, title=title, metadata=metadata or {}
|
url, title=title, metadata=metadata or {}
|
||||||
)
|
)
|
||||||
return document.id
|
# Handle both single document and list of documents
|
||||||
|
if isinstance(result, list):
|
||||||
|
return result[0].id if result else None
|
||||||
|
return result.id
|
||||||
except Exception:
|
except Exception:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
@ -188,8 +194,8 @@ def create_mcp_server(db_path: Path) -> FastMCP:
|
||||||
deps = DeepQADeps(client=rag)
|
deps = DeepQADeps(client=rag)
|
||||||
|
|
||||||
start_node = DeepQAPlanNode(
|
start_node = DeepQAPlanNode(
|
||||||
provider=Config.QA_PROVIDER,
|
provider=Config.qa.provider,
|
||||||
model=Config.QA_MODEL,
|
model=Config.qa.model,
|
||||||
)
|
)
|
||||||
|
|
||||||
result = await graph.run(
|
result = await graph.run(
|
||||||
|
|
@ -241,8 +247,8 @@ def create_mcp_server(db_path: Path) -> FastMCP:
|
||||||
|
|
||||||
result = await graph.run(
|
result = await graph.run(
|
||||||
PlanNode(
|
PlanNode(
|
||||||
provider=Config.RESEARCH_PROVIDER or Config.QA_PROVIDER,
|
provider=Config.research.provider or Config.qa.provider,
|
||||||
model=Config.RESEARCH_MODEL or Config.QA_MODEL,
|
model=Config.research.model or Config.qa.model,
|
||||||
),
|
),
|
||||||
state=state,
|
state=state,
|
||||||
deps=deps,
|
deps=deps,
|
||||||
|
|
|
||||||
|
|
@ -8,7 +8,7 @@ from haiku.rag.client import HaikuRAG
|
||||||
from haiku.rag.store.models.document import Document
|
from haiku.rag.store.models.document import Document
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from haiku.rag.reader import FileReader
|
pass
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
@ -46,6 +46,9 @@ class FileWatcher:
|
||||||
await self._delete_document(Path(path))
|
await self._delete_document(Path(path))
|
||||||
|
|
||||||
async def refresh(self):
|
async def refresh(self):
|
||||||
|
# Lazy import to avoid loading docling
|
||||||
|
from haiku.rag.reader import FileReader
|
||||||
|
|
||||||
for path in self.paths:
|
for path in self.paths:
|
||||||
for f in Path(path).rglob("**/*"):
|
for f in Path(path).rglob("**/*"):
|
||||||
if f.is_file() and f.suffix in FileReader.extensions:
|
if f.is_file() and f.suffix in FileReader.extensions:
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
|
|
@ -71,13 +71,15 @@ class QuestionAnswerAgent:
|
||||||
if provider == "ollama":
|
if provider == "ollama":
|
||||||
return OpenAIChatModel(
|
return OpenAIChatModel(
|
||||||
model_name=model,
|
model_name=model,
|
||||||
provider=OllamaProvider(base_url=f"{Config.OLLAMA_BASE_URL}/v1"),
|
provider=OllamaProvider(
|
||||||
|
base_url=f"{Config.providers.ollama.base_url}/v1"
|
||||||
|
),
|
||||||
)
|
)
|
||||||
elif provider == "vllm":
|
elif provider == "vllm":
|
||||||
return OpenAIChatModel(
|
return OpenAIChatModel(
|
||||||
model_name=model,
|
model_name=model,
|
||||||
provider=OpenAIProvider(
|
provider=OpenAIProvider(
|
||||||
base_url=f"{Config.VLLM_QA_BASE_URL}/v1", api_key="none"
|
base_url=f"{Config.providers.vllm.qa_base_url}/v1", api_key="none"
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
|
|
||||||
|
|
@ -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.RERANK_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.RERANK_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
|
||||||
|
|
|
||||||
|
|
@ -3,7 +3,7 @@ from haiku.rag.store.models.chunk import Chunk
|
||||||
|
|
||||||
|
|
||||||
class RerankerBase:
|
class RerankerBase:
|
||||||
_model: str = Config.RERANK_MODEL
|
_model: str = Config.reranking.model
|
||||||
|
|
||||||
async def rerank(
|
async def rerank(
|
||||||
self, query: str, chunks: list[Chunk], top_n: int = 10
|
self, query: str, chunks: list[Chunk], top_n: int = 10
|
||||||
|
|
|
||||||
|
|
@ -1,4 +1,3 @@
|
||||||
from haiku.rag.config import Config
|
|
||||||
from haiku.rag.reranking.base import RerankerBase
|
from haiku.rag.reranking.base import RerankerBase
|
||||||
from haiku.rag.store.models.chunk import Chunk
|
from haiku.rag.store.models.chunk import Chunk
|
||||||
|
|
||||||
|
|
@ -12,7 +11,8 @@ except ImportError as e:
|
||||||
|
|
||||||
class CohereReranker(RerankerBase):
|
class CohereReranker(RerankerBase):
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self._client = cohere.ClientV2(api_key=Config.COHERE_API_KEY)
|
# Cohere SDK reads CO_API_KEY from environment by default
|
||||||
|
self._client = cohere.ClientV2()
|
||||||
|
|
||||||
async def rerank(
|
async def rerank(
|
||||||
self, query: str, chunks: list[Chunk], top_n: int = 10
|
self, query: str, chunks: list[Chunk], top_n: int = 10
|
||||||
|
|
|
||||||
|
|
@ -8,7 +8,7 @@ from haiku.rag.store.models.chunk import Chunk
|
||||||
class MxBAIReranker(RerankerBase):
|
class MxBAIReranker(RerankerBase):
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self._client = MxbaiRerankV2(
|
self._client = MxbaiRerankV2(
|
||||||
Config.RERANK_MODEL, disable_transformers_warnings=True
|
Config.reranking.model, disable_transformers_warnings=True
|
||||||
)
|
)
|
||||||
|
|
||||||
async def rerank(
|
async def rerank(
|
||||||
|
|
|
||||||
|
|
@ -8,7 +8,7 @@ from haiku.rag.store.models.chunk import Chunk
|
||||||
class VLLMReranker(RerankerBase):
|
class VLLMReranker(RerankerBase):
|
||||||
def __init__(self, model: str):
|
def __init__(self, model: str):
|
||||||
self._model = model
|
self._model = model
|
||||||
self._base_url = Config.VLLM_RERANK_BASE_URL
|
self._base_url = Config.providers.vllm.rerank_base_url
|
||||||
|
|
||||||
async def rerank(
|
async def rerank(
|
||||||
self, query: str, chunks: list[Chunk], top_n: int = 10
|
self, query: str, chunks: list[Chunk], top_n: int = 10
|
||||||
|
|
|
||||||
|
|
@ -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.DISABLE_DB_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.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.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))]
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -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.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)
|
||||||
|
|
|
||||||
|
|
@ -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,17 +115,28 @@ 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
|
||||||
stored_provider = stored_settings.get("EMBEDDINGS_PROVIDER")
|
# Support both old flat structure and new nested structure for backward compatibility
|
||||||
current_provider = current_config.get("EMBEDDINGS_PROVIDER")
|
stored_embeddings = stored_settings.get("embeddings", {})
|
||||||
|
current_embeddings = current_config.get("embeddings", {})
|
||||||
|
|
||||||
stored_model = stored_settings.get("EMBEDDINGS_MODEL")
|
# Try nested structure first, fall back to flat for old databases
|
||||||
current_model = current_config.get("EMBEDDINGS_MODEL")
|
stored_provider = stored_embeddings.get("provider") or stored_settings.get(
|
||||||
|
"EMBEDDINGS_PROVIDER"
|
||||||
|
)
|
||||||
|
current_provider = current_embeddings.get("provider")
|
||||||
|
|
||||||
stored_vector_dim = stored_settings.get("EMBEDDINGS_VECTOR_DIM")
|
stored_model = stored_embeddings.get("model") or stored_settings.get(
|
||||||
current_vector_dim = current_config.get("EMBEDDINGS_VECTOR_DIM")
|
"EMBEDDINGS_MODEL"
|
||||||
|
)
|
||||||
|
current_model = current_embeddings.get("model")
|
||||||
|
|
||||||
|
stored_vector_dim = stored_embeddings.get("vector_dim") or stored_settings.get(
|
||||||
|
"EMBEDDINGS_VECTOR_DIM"
|
||||||
|
)
|
||||||
|
current_vector_dim = current_embeddings.get("vector_dim")
|
||||||
|
|
||||||
# Check for incompatible changes
|
# Check for incompatible changes
|
||||||
incompatible_changes = []
|
incompatible_changes = []
|
||||||
|
|
|
||||||
|
|
@ -176,19 +176,19 @@ def prefetch_models():
|
||||||
|
|
||||||
# Collect Ollama models from config
|
# Collect Ollama models from config
|
||||||
required_models: set[str] = set()
|
required_models: set[str] = set()
|
||||||
if Config.EMBEDDINGS_PROVIDER == "ollama":
|
if Config.embeddings.provider == "ollama":
|
||||||
required_models.add(Config.EMBEDDINGS_MODEL)
|
required_models.add(Config.embeddings.model)
|
||||||
if Config.QA_PROVIDER == "ollama":
|
if Config.qa.provider == "ollama":
|
||||||
required_models.add(Config.QA_MODEL)
|
required_models.add(Config.qa.model)
|
||||||
if Config.RESEARCH_PROVIDER == "ollama":
|
if Config.research.provider == "ollama":
|
||||||
required_models.add(Config.RESEARCH_MODEL)
|
required_models.add(Config.research.model)
|
||||||
if Config.RERANK_PROVIDER == "ollama":
|
if Config.reranking.provider == "ollama":
|
||||||
required_models.add(Config.RERANK_MODEL)
|
required_models.add(Config.reranking.model)
|
||||||
|
|
||||||
if not required_models:
|
if not required_models:
|
||||||
return
|
return
|
||||||
|
|
||||||
base_url = Config.OLLAMA_BASE_URL
|
base_url = Config.providers.ollama.base_url
|
||||||
|
|
||||||
with httpx.Client(timeout=None) as client:
|
with httpx.Client(timeout=None) as client:
|
||||||
for model in sorted(required_models):
|
for model in sorted(required_models):
|
||||||
|
|
|
||||||
|
|
@ -2,6 +2,7 @@ import tempfile
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
import yaml
|
||||||
from datasets import Dataset, load_dataset, load_from_disk
|
from datasets import Dataset, load_dataset, load_from_disk
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -24,3 +25,36 @@ def temp_db_path():
|
||||||
"""Create a temporary database path for testing."""
|
"""Create a temporary database path for testing."""
|
||||||
with tempfile.TemporaryDirectory() as temp_dir:
|
with tempfile.TemporaryDirectory() as temp_dir:
|
||||||
yield Path(temp_dir) / "test.lancedb"
|
yield Path(temp_dir) / "test.lancedb"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def temp_yaml_config(tmp_path, monkeypatch):
|
||||||
|
"""Create a temporary YAML config file for testing.
|
||||||
|
|
||||||
|
This fixture creates a config file in a temp directory and sets
|
||||||
|
the environment variable so config.py will load it.
|
||||||
|
"""
|
||||||
|
config_file = tmp_path / "test-config.yaml"
|
||||||
|
config_data = {
|
||||||
|
"environment": "production",
|
||||||
|
"storage": {
|
||||||
|
"data_dir": "",
|
||||||
|
"monitor_directories": [],
|
||||||
|
"disable_autocreate": False,
|
||||||
|
"vacuum_retention_seconds": 60,
|
||||||
|
},
|
||||||
|
"embeddings": {
|
||||||
|
"provider": "ollama",
|
||||||
|
"model": "qwen3-embedding",
|
||||||
|
"vector_dim": 4096,
|
||||||
|
},
|
||||||
|
"qa": {"provider": "ollama", "model": "gpt-oss"},
|
||||||
|
}
|
||||||
|
|
||||||
|
with open(config_file, "w") as f:
|
||||||
|
yaml.dump(config_data, f)
|
||||||
|
|
||||||
|
# Set env var so config loader will find it
|
||||||
|
monkeypatch.setenv("HAIKU_RAG_CONFIG_PATH", str(config_file))
|
||||||
|
|
||||||
|
yield config_file
|
||||||
|
|
|
||||||
|
|
@ -557,7 +557,7 @@ async def test_client_create_document_with_custom_chunks(temp_db_path):
|
||||||
Chunk(
|
Chunk(
|
||||||
content="This is the second chunk",
|
content="This is the second chunk",
|
||||||
metadata={"custom": "metadata2"},
|
metadata={"custom": "metadata2"},
|
||||||
embedding=[0.1] * Config.EMBEDDINGS_VECTOR_DIM,
|
embedding=[0.1] * Config.embeddings.vector_dim,
|
||||||
order=1,
|
order=1,
|
||||||
), # With embedding
|
), # With embedding
|
||||||
Chunk(
|
Chunk(
|
||||||
|
|
@ -641,7 +641,7 @@ async def test_client_ask_with_cite(monkeypatch, temp_db_path):
|
||||||
async def test_client_expand_context(temp_db_path):
|
async def test_client_expand_context(temp_db_path):
|
||||||
"""Test expanding search results with adjacent chunks."""
|
"""Test expanding search results with adjacent chunks."""
|
||||||
# Mock Config to have CONTEXT_CHUNK_RADIUS = 2
|
# Mock Config to have CONTEXT_CHUNK_RADIUS = 2
|
||||||
with patch("haiku.rag.client.Config.CONTEXT_CHUNK_RADIUS", 2):
|
with patch("haiku.rag.client.Config.processing.context_chunk_radius", 2):
|
||||||
async with HaikuRAG(temp_db_path) as client:
|
async with HaikuRAG(temp_db_path) as client:
|
||||||
# Create chunks manually with precomputed embeddings to avoid network
|
# Create chunks manually with precomputed embeddings to avoid network
|
||||||
dim = client.chunk_repository.embedder._vector_dim
|
dim = client.chunk_repository.embedder._vector_dim
|
||||||
|
|
@ -710,7 +710,7 @@ async def test_client_expand_context_radius_zero(temp_db_path):
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_client_expand_context_multiple_chunks(temp_db_path):
|
async def test_client_expand_context_multiple_chunks(temp_db_path):
|
||||||
"""Test expand_context with multiple search results."""
|
"""Test expand_context with multiple search results."""
|
||||||
with patch("haiku.rag.client.Config.CONTEXT_CHUNK_RADIUS", 1):
|
with patch("haiku.rag.client.Config.processing.context_chunk_radius", 1):
|
||||||
async with HaikuRAG(temp_db_path) as client:
|
async with HaikuRAG(temp_db_path) as client:
|
||||||
# Create first document with manual chunks
|
# Create first document with manual chunks
|
||||||
doc1_chunks = [
|
doc1_chunks = [
|
||||||
|
|
|
||||||
192
tests/test_config_loader.py
Normal file
192
tests/test_config_loader.py
Normal file
|
|
@ -0,0 +1,192 @@
|
||||||
|
import os
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from haiku.rag.config.loader import (
|
||||||
|
find_config_file,
|
||||||
|
generate_default_config,
|
||||||
|
load_config_from_env,
|
||||||
|
load_yaml_config,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_load_yaml_config(tmp_path):
|
||||||
|
"""Test loading a YAML config file."""
|
||||||
|
config_file = tmp_path / "test.yaml"
|
||||||
|
config_file.write_text("""
|
||||||
|
environment: production
|
||||||
|
embeddings:
|
||||||
|
provider: ollama
|
||||||
|
model: test-model
|
||||||
|
vector_dim: 1024
|
||||||
|
""")
|
||||||
|
|
||||||
|
config = load_yaml_config(config_file)
|
||||||
|
assert config["environment"] == "production"
|
||||||
|
assert config["embeddings"]["provider"] == "ollama"
|
||||||
|
assert config["embeddings"]["model"] == "test-model"
|
||||||
|
assert config["embeddings"]["vector_dim"] == 1024
|
||||||
|
|
||||||
|
|
||||||
|
def test_find_config_file_cwd(tmp_path, monkeypatch):
|
||||||
|
"""Test finding config in current directory."""
|
||||||
|
monkeypatch.chdir(tmp_path)
|
||||||
|
config_file = tmp_path / "haiku.rag.yaml"
|
||||||
|
config_file.write_text("environment: production")
|
||||||
|
|
||||||
|
found = find_config_file()
|
||||||
|
assert found == config_file
|
||||||
|
|
||||||
|
|
||||||
|
def test_find_config_file_user_config(tmp_path, monkeypatch):
|
||||||
|
"""Test finding config in user config directory."""
|
||||||
|
monkeypatch.chdir(tmp_path)
|
||||||
|
user_config_dir = tmp_path / ".config" / "haiku.rag"
|
||||||
|
user_config_dir.mkdir(parents=True)
|
||||||
|
config_file = user_config_dir / "config.yaml"
|
||||||
|
config_file.write_text("environment: production")
|
||||||
|
|
||||||
|
# Mock home directory
|
||||||
|
monkeypatch.setattr(Path, "home", lambda: tmp_path)
|
||||||
|
|
||||||
|
found = find_config_file()
|
||||||
|
assert found == config_file
|
||||||
|
|
||||||
|
|
||||||
|
def test_find_config_file_cli_path(tmp_path):
|
||||||
|
"""Test finding config via CLI path parameter."""
|
||||||
|
config_file = tmp_path / "custom.yaml"
|
||||||
|
config_file.write_text("environment: production")
|
||||||
|
|
||||||
|
found = find_config_file(config_file)
|
||||||
|
assert found == config_file
|
||||||
|
|
||||||
|
|
||||||
|
def test_find_config_file_env_var(tmp_path, monkeypatch):
|
||||||
|
"""Test finding config via HAIKU_RAG_CONFIG_PATH env var."""
|
||||||
|
config_file = tmp_path / "from-env.yaml"
|
||||||
|
config_file.write_text("environment: production")
|
||||||
|
|
||||||
|
monkeypatch.setenv("HAIKU_RAG_CONFIG_PATH", str(config_file))
|
||||||
|
|
||||||
|
found = find_config_file()
|
||||||
|
assert found == config_file
|
||||||
|
|
||||||
|
|
||||||
|
def test_find_config_file_not_found(tmp_path, monkeypatch):
|
||||||
|
"""Test returning None when no config found."""
|
||||||
|
monkeypatch.chdir(tmp_path)
|
||||||
|
monkeypatch.setattr(Path, "home", lambda: tmp_path)
|
||||||
|
|
||||||
|
found = find_config_file()
|
||||||
|
assert found is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_find_config_file_cli_path_not_exists(tmp_path):
|
||||||
|
"""Test error when CLI path doesn't exist."""
|
||||||
|
config_file = tmp_path / "nonexistent.yaml"
|
||||||
|
|
||||||
|
with pytest.raises(FileNotFoundError):
|
||||||
|
find_config_file(config_file)
|
||||||
|
|
||||||
|
|
||||||
|
def test_generate_default_config():
|
||||||
|
"""Test generating default config structure."""
|
||||||
|
config = generate_default_config()
|
||||||
|
|
||||||
|
assert config["environment"] == "production"
|
||||||
|
assert "storage" in config
|
||||||
|
assert "embeddings" in config
|
||||||
|
assert "qa" in config
|
||||||
|
assert "providers" in config
|
||||||
|
assert config["embeddings"]["provider"] == "ollama"
|
||||||
|
assert config["embeddings"]["vector_dim"] == 4096
|
||||||
|
|
||||||
|
|
||||||
|
def test_load_config_from_env(monkeypatch):
|
||||||
|
"""Test loading config from environment variables."""
|
||||||
|
monkeypatch.setenv("ENV", "development")
|
||||||
|
monkeypatch.setenv("EMBEDDINGS_PROVIDER", "openai")
|
||||||
|
monkeypatch.setenv("EMBEDDINGS_MODEL", "text-embedding-3-small")
|
||||||
|
monkeypatch.setenv("EMBEDDINGS_VECTOR_DIM", "1536")
|
||||||
|
monkeypatch.setenv("QA_PROVIDER", "anthropic")
|
||||||
|
monkeypatch.setenv("QA_MODEL", "claude-3-haiku")
|
||||||
|
|
||||||
|
config = load_config_from_env()
|
||||||
|
|
||||||
|
assert config["environment"] == "development"
|
||||||
|
assert config["embeddings"]["provider"] == "openai"
|
||||||
|
assert config["embeddings"]["model"] == "text-embedding-3-small"
|
||||||
|
assert config["embeddings"]["vector_dim"] == "1536"
|
||||||
|
assert config["qa"]["provider"] == "anthropic"
|
||||||
|
assert config["qa"]["model"] == "claude-3-haiku"
|
||||||
|
|
||||||
|
|
||||||
|
def test_load_config_from_env_empty():
|
||||||
|
"""Test loading from env when no relevant vars set."""
|
||||||
|
# Clear any env vars that might be set
|
||||||
|
env_vars = [
|
||||||
|
"ENV",
|
||||||
|
"EMBEDDINGS_PROVIDER",
|
||||||
|
"QA_PROVIDER",
|
||||||
|
"OPENAI_API_KEY",
|
||||||
|
]
|
||||||
|
original_values = {}
|
||||||
|
for var in env_vars:
|
||||||
|
original_values[var] = os.environ.get(var)
|
||||||
|
if var in os.environ:
|
||||||
|
del os.environ[var]
|
||||||
|
|
||||||
|
try:
|
||||||
|
config = load_config_from_env()
|
||||||
|
# Should return empty or minimal dict
|
||||||
|
assert isinstance(config, dict)
|
||||||
|
finally:
|
||||||
|
# Restore original values
|
||||||
|
for var, value in original_values.items():
|
||||||
|
if value is not None:
|
||||||
|
os.environ[var] = value
|
||||||
|
|
||||||
|
|
||||||
|
def test_config_precedence_cwd_over_user(tmp_path, monkeypatch):
|
||||||
|
"""Test that cwd config takes precedence over user config."""
|
||||||
|
monkeypatch.chdir(tmp_path)
|
||||||
|
monkeypatch.setattr(Path, "home", lambda: tmp_path)
|
||||||
|
|
||||||
|
# Create both configs
|
||||||
|
cwd_config = tmp_path / "haiku.rag.yaml"
|
||||||
|
cwd_config.write_text("environment: from-cwd")
|
||||||
|
|
||||||
|
user_config_dir = tmp_path / ".config" / "haiku.rag"
|
||||||
|
user_config_dir.mkdir(parents=True)
|
||||||
|
user_config = user_config_dir / "config.yaml"
|
||||||
|
user_config.write_text("environment: from-user")
|
||||||
|
|
||||||
|
found = find_config_file()
|
||||||
|
assert found == cwd_config
|
||||||
|
assert found is not None
|
||||||
|
|
||||||
|
config = load_yaml_config(found)
|
||||||
|
assert config["environment"] == "from-cwd"
|
||||||
|
|
||||||
|
|
||||||
|
def test_config_precedence_env_var_over_cwd(tmp_path, monkeypatch):
|
||||||
|
"""Test that HAIKU_RAG_CONFIG_PATH env var takes precedence."""
|
||||||
|
monkeypatch.chdir(tmp_path)
|
||||||
|
|
||||||
|
# Create cwd config
|
||||||
|
cwd_config = tmp_path / "haiku.rag.yaml"
|
||||||
|
cwd_config.write_text("environment: from-cwd")
|
||||||
|
|
||||||
|
# Create env var config
|
||||||
|
env_config = tmp_path / "from-env.yaml"
|
||||||
|
env_config.write_text("environment: from-env-var")
|
||||||
|
monkeypatch.setenv("HAIKU_RAG_CONFIG_PATH", str(env_config))
|
||||||
|
|
||||||
|
found = find_config_file()
|
||||||
|
assert found == env_config
|
||||||
|
assert found is not None
|
||||||
|
|
||||||
|
config = load_yaml_config(found)
|
||||||
|
assert config["environment"] == "from-env-var"
|
||||||
|
|
@ -1,3 +1,5 @@
|
||||||
|
import os
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
|
|
@ -6,9 +8,9 @@ from haiku.rag.embeddings.ollama import Embedder as OllamaEmbedder
|
||||||
from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder
|
from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder
|
||||||
from haiku.rag.embeddings.vllm import Embedder as VLLMEmbedder
|
from haiku.rag.embeddings.vllm import Embedder as VLLMEmbedder
|
||||||
|
|
||||||
OPENAI_AVAILABLE = bool(Config.OPENAI_API_KEY)
|
OPENAI_AVAILABLE = bool(os.getenv("OPENAI_API_KEY"))
|
||||||
VOYAGEAI_AVAILABLE = bool(Config.VOYAGE_API_KEY)
|
VOYAGEAI_AVAILABLE = bool(os.getenv("VOYAGE_API_KEY"))
|
||||||
VLLM_EMBEDDINGS_AVAILABLE = bool(Config.VLLM_EMBEDDINGS_BASE_URL)
|
VLLM_EMBEDDINGS_AVAILABLE = bool(Config.providers.vllm.embeddings_base_url)
|
||||||
|
|
||||||
|
|
||||||
# Calculate cosine similarity
|
# Calculate cosine similarity
|
||||||
|
|
|
||||||
|
|
@ -14,9 +14,9 @@ async def test_lancedb_cloud_skips_optimization(temp_db_path):
|
||||||
|
|
||||||
# Mock all cloud config to simulate LanceDB Cloud usage
|
# Mock all cloud config to simulate LanceDB Cloud usage
|
||||||
with (
|
with (
|
||||||
patch.object(Config, "LANCEDB_URI", "db://test-database"),
|
patch.object(Config.lancedb, "uri", "db://test-database"),
|
||||||
patch.object(Config, "LANCEDB_API_KEY", "test-api-key"),
|
patch.object(Config.lancedb, "api_key", "test-api-key"),
|
||||||
patch.object(Config, "LANCEDB_REGION", "us-east-1"),
|
patch.object(Config.lancedb, "region", "us-east-1"),
|
||||||
):
|
):
|
||||||
# Mock the optimize method to track if it's called
|
# Mock the optimize method to track if it's called
|
||||||
with patch.object(store.chunks_table, "optimize") as mock_optimize:
|
with patch.object(store.chunks_table, "optimize") as mock_optimize:
|
||||||
|
|
@ -35,8 +35,8 @@ async def test_local_storage_calls_optimization(temp_db_path):
|
||||||
# Create a store
|
# Create a store
|
||||||
store = Store(temp_db_path)
|
store = Store(temp_db_path)
|
||||||
|
|
||||||
# Ensure LANCEDB_URI is empty (local storage)
|
# Ensure uri is empty (local storage)
|
||||||
with patch.object(Config, "LANCEDB_URI", ""):
|
with patch.object(Config.lancedb, "uri", ""):
|
||||||
# Mock the optimize method to track if it's called
|
# Mock the optimize method to track if it's called
|
||||||
with patch.object(store.chunks_table, "optimize") as mock_optimize:
|
with patch.object(store.chunks_table, "optimize") as mock_optimize:
|
||||||
# Call vacuum - this should optimize all tables for local storage
|
# Call vacuum - this should optimize all tables for local storage
|
||||||
|
|
|
||||||
|
|
@ -41,9 +41,9 @@ def add_marker(text: str) -> str:
|
||||||
"""
|
"""
|
||||||
)
|
)
|
||||||
|
|
||||||
original_pre = Config.MARKDOWN_PREPROCESSOR
|
original_pre = Config.processing.markdown_preprocessor
|
||||||
try:
|
try:
|
||||||
Config.MARKDOWN_PREPROCESSOR = f"{pre_file}:add_marker"
|
Config.processing.markdown_preprocessor = f"{pre_file}:add_marker"
|
||||||
|
|
||||||
store = Store(temp_db_path)
|
store = Store(temp_db_path)
|
||||||
chunk_repo = ChunkRepository(store)
|
chunk_repo = ChunkRepository(store)
|
||||||
|
|
@ -68,4 +68,4 @@ def add_marker(text: str) -> str:
|
||||||
|
|
||||||
assert any(marker in c.content for c in chunks)
|
assert any(marker in c.content for c in chunks)
|
||||||
finally:
|
finally:
|
||||||
Config.MARKDOWN_PREPROCESSOR = original_pre
|
Config.processing.markdown_preprocessor = original_pre
|
||||||
|
|
|
||||||
|
|
@ -1,3 +1,5 @@
|
||||||
|
import os
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from datasets import Dataset
|
from datasets import Dataset
|
||||||
|
|
||||||
|
|
@ -6,9 +8,9 @@ from haiku.rag.client import HaikuRAG
|
||||||
from haiku.rag.config import Config
|
from haiku.rag.config import Config
|
||||||
from haiku.rag.qa.agent import QuestionAnswerAgent
|
from haiku.rag.qa.agent import QuestionAnswerAgent
|
||||||
|
|
||||||
OPENAI_AVAILABLE = bool(Config.OPENAI_API_KEY)
|
OPENAI_AVAILABLE = bool(os.getenv("OPENAI_API_KEY"))
|
||||||
ANTHROPIC_AVAILABLE = bool(Config.ANTHROPIC_API_KEY)
|
ANTHROPIC_AVAILABLE = bool(os.getenv("ANTHROPIC_API_KEY"))
|
||||||
VLLM_QA_AVAILABLE = bool(Config.VLLM_QA_BASE_URL)
|
VLLM_QA_AVAILABLE = bool(Config.providers.vllm.qa_base_url)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|
|
||||||
|
|
@ -1,3 +1,5 @@
|
||||||
|
import os
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from haiku.rag.config import Config
|
from haiku.rag.config import Config
|
||||||
|
|
@ -5,8 +7,8 @@ from haiku.rag.reranking.base import RerankerBase
|
||||||
from haiku.rag.reranking.vllm import VLLMReranker
|
from haiku.rag.reranking.vllm import VLLMReranker
|
||||||
from haiku.rag.store.models.chunk import Chunk
|
from haiku.rag.store.models.chunk import Chunk
|
||||||
|
|
||||||
COHERE_AVAILABLE = bool(Config.COHERE_API_KEY)
|
COHERE_AVAILABLE = bool(os.getenv("CO_API_KEY"))
|
||||||
VLLM_RERANK_AVAILABLE = bool(Config.VLLM_RERANK_BASE_URL)
|
VLLM_RERANK_AVAILABLE = bool(Config.providers.vllm.rerank_base_url)
|
||||||
|
|
||||||
chunks = [
|
chunks = [
|
||||||
Chunk(content=content, document_id=str(i))
|
Chunk(content=content, document_id=str(i))
|
||||||
|
|
@ -37,7 +39,7 @@ async def test_mxbai_reranker():
|
||||||
try:
|
try:
|
||||||
from haiku.rag.reranking.mxbai import MxBAIReranker
|
from haiku.rag.reranking.mxbai import MxBAIReranker
|
||||||
|
|
||||||
Config.RERANK_MODEL = "mixedbread-ai/mxbai-rerank-base-v2"
|
Config.reranking.model = "mixedbread-ai/mxbai-rerank-base-v2"
|
||||||
reranker = MxBAIReranker()
|
reranker = MxBAIReranker()
|
||||||
# reranker._model = "mixedbread-ai/mxbai-rerank-base-v2"
|
# reranker._model = "mixedbread-ai/mxbai-rerank-base-v2"
|
||||||
reranked = await reranker.rerank(
|
reranked = await reranker.rerank(
|
||||||
|
|
@ -45,7 +47,7 @@ async def test_mxbai_reranker():
|
||||||
)
|
)
|
||||||
assert [chunk.document_id for chunk, score in reranked] == ["0", "2"]
|
assert [chunk.document_id for chunk, score in reranked] == ["0", "2"]
|
||||||
assert all(isinstance(score, float) for chunk, score in reranked)
|
assert all(isinstance(score, float) for chunk, score in reranked)
|
||||||
Config.RERANK_MODEL = ""
|
Config.reranking.model = ""
|
||||||
|
|
||||||
except ImportError:
|
except ImportError:
|
||||||
pytest.skip("MxBAI package not installed")
|
pytest.skip("MxBAI package not installed")
|
||||||
|
|
|
||||||
|
|
@ -35,14 +35,14 @@ def test_settings_save_and_retrieve(temp_db_path):
|
||||||
store = Store(temp_db_path)
|
store = Store(temp_db_path)
|
||||||
settings_repo = SettingsRepository(store)
|
settings_repo = SettingsRepository(store)
|
||||||
|
|
||||||
original_chunk_size = Config.CHUNK_SIZE
|
original_chunk_size = Config.processing.chunk_size
|
||||||
Config.CHUNK_SIZE = 2 * original_chunk_size
|
Config.processing.chunk_size = 2 * original_chunk_size
|
||||||
|
|
||||||
settings_repo.save_current_settings()
|
settings_repo.save_current_settings()
|
||||||
retrieved_settings = settings_repo.get_current_settings()
|
retrieved_settings = settings_repo.get_current_settings()
|
||||||
assert retrieved_settings["CHUNK_SIZE"] == 2 * original_chunk_size
|
assert retrieved_settings["processing"]["chunk_size"] == 2 * original_chunk_size
|
||||||
|
|
||||||
Config.CHUNK_SIZE = original_chunk_size
|
Config.processing.chunk_size = original_chunk_size
|
||||||
store.close()
|
store.close()
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -57,16 +57,16 @@ async def test_config_validation_on_db_load(temp_db_path):
|
||||||
store1.close()
|
store1.close()
|
||||||
|
|
||||||
# Change config
|
# Change config
|
||||||
original_chunk_size = Config.CHUNK_SIZE
|
original_chunk_size = Config.processing.chunk_size
|
||||||
Config.CHUNK_SIZE = 999
|
Config.processing.chunk_size = 999
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# Loading the database should raise ConfigMismatchError
|
# Loading the database should raise ConfigMismatchError
|
||||||
with pytest.raises(ConfigMismatchError) as exc_info:
|
with pytest.raises(ConfigMismatchError) as exc_info:
|
||||||
Store(temp_db_path)
|
Store(temp_db_path)
|
||||||
|
|
||||||
assert "CHUNK_SIZE" in str(exc_info.value)
|
assert "chunk_size" in str(exc_info.value)
|
||||||
assert "Consider rebuilding" in str(exc_info.value)
|
assert "rebuild" in str(exc_info.value).lower()
|
||||||
|
|
||||||
# Rebuild
|
# Rebuild
|
||||||
async with HaikuRAG(db_path=temp_db_path, skip_validation=True) as client:
|
async with HaikuRAG(db_path=temp_db_path, skip_validation=True) as client:
|
||||||
|
|
@ -77,8 +77,8 @@ async def test_config_validation_on_db_load(temp_db_path):
|
||||||
store2 = Store(temp_db_path)
|
store2 = Store(temp_db_path)
|
||||||
settings_repo2 = SettingsRepository(store2)
|
settings_repo2 = SettingsRepository(store2)
|
||||||
db_settings = settings_repo2.get_current_settings()
|
db_settings = settings_repo2.get_current_settings()
|
||||||
assert db_settings["CHUNK_SIZE"] == 999
|
assert db_settings["processing"]["chunk_size"] == 999
|
||||||
store2.close()
|
store2.close()
|
||||||
|
|
||||||
finally:
|
finally:
|
||||||
Config.CHUNK_SIZE = original_chunk_size
|
Config.processing.chunk_size = original_chunk_size
|
||||||
|
|
|
||||||
|
|
@ -204,7 +204,7 @@ async def test_vacuum_completes_before_context_exit(temp_db_path, monkeypatch):
|
||||||
from haiku.rag.utils import text_to_docling_document
|
from haiku.rag.utils import text_to_docling_document
|
||||||
|
|
||||||
# Set aggressive vacuum retention for this test
|
# Set aggressive vacuum retention for this test
|
||||||
monkeypatch.setattr(Config, "VACUUM_RETENTION_SECONDS", 0)
|
monkeypatch.setattr(Config.storage, "vacuum_retention_seconds", 0)
|
||||||
|
|
||||||
async with HaikuRAG(db_path=temp_db_path) as client:
|
async with HaikuRAG(db_path=temp_db_path) as client:
|
||||||
# Create multiple documents - each creation triggers automatic vacuum with retention=0
|
# Create multiple documents - each creation triggers automatic vacuum with retention=0
|
||||||
|
|
|
||||||
2
uv.lock
2
uv.lock
|
|
@ -1134,6 +1134,7 @@ dependencies = [
|
||||||
{ name = "pydantic-ai" },
|
{ name = "pydantic-ai" },
|
||||||
{ name = "pydantic-graph" },
|
{ name = "pydantic-graph" },
|
||||||
{ name = "python-dotenv" },
|
{ name = "python-dotenv" },
|
||||||
|
{ name = "pyyaml" },
|
||||||
{ name = "rich" },
|
{ name = "rich" },
|
||||||
{ name = "tiktoken" },
|
{ name = "tiktoken" },
|
||||||
{ name = "typer" },
|
{ name = "typer" },
|
||||||
|
|
@ -1178,6 +1179,7 @@ requires-dist = [
|
||||||
{ name = "pydantic-ai", specifier = ">=1.0.18" },
|
{ name = "pydantic-ai", specifier = ">=1.0.18" },
|
||||||
{ name = "pydantic-graph", specifier = ">=1.0.18" },
|
{ name = "pydantic-graph", specifier = ">=1.0.18" },
|
||||||
{ name = "python-dotenv", specifier = ">=1.1.1" },
|
{ name = "python-dotenv", specifier = ">=1.1.1" },
|
||||||
|
{ name = "pyyaml", specifier = ">=6.0.1" },
|
||||||
{ name = "rich", specifier = ">=14.2.0" },
|
{ name = "rich", specifier = ">=14.2.0" },
|
||||||
{ name = "tiktoken", specifier = ">=0.12.0" },
|
{ name = "tiktoken", specifier = ">=0.12.0" },
|
||||||
{ name = "typer", specifier = ">=0.19.2" },
|
{ name = "typer", specifier = ">=0.19.2" },
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue