Use EmbeddingModelConfig similar to ModelConfig for embeddings
This commit is contained in:
parent
f07590bd5b
commit
3525fae625
19 changed files with 165 additions and 89 deletions
|
|
@ -41,9 +41,9 @@ def build_experiment_metadata(
|
||||||
return {
|
return {
|
||||||
"dataset": dataset_key,
|
"dataset": dataset_key,
|
||||||
"test_cases": test_cases,
|
"test_cases": test_cases,
|
||||||
"embedder_provider": config.embeddings.provider,
|
"embedder_provider": config.embeddings.model.provider,
|
||||||
"embedder_model": config.embeddings.model,
|
"embedder_model": config.embeddings.model.name,
|
||||||
"embedder_dim": config.embeddings.vector_dim,
|
"embedder_dim": config.embeddings.model.vector_dim,
|
||||||
"chunk_size": config.processing.chunk_size,
|
"chunk_size": config.processing.chunk_size,
|
||||||
"context_chunk_radius": config.processing.context_chunk_radius,
|
"context_chunk_radius": config.processing.context_chunk_radius,
|
||||||
"rerank_provider": config.reranking.model.provider
|
"rerank_provider": config.reranking.model.provider
|
||||||
|
|
|
||||||
|
|
@ -38,9 +38,7 @@ class LLMJudge:
|
||||||
|
|
||||||
def __init__(self, model: str = "gpt-oss"):
|
def __init__(self, model: str = "gpt-oss"):
|
||||||
# Create model using get_model with thinking disabled
|
# Create model using get_model with thinking disabled
|
||||||
model_config = ModelConfig(
|
model_config = ModelConfig(provider="ollama", name=model, enable_thinking=False)
|
||||||
provider="ollama", model=model, enable_thinking=False
|
|
||||||
)
|
|
||||||
model_obj = get_model(model_config, Config)
|
model_obj = get_model(model_config, Config)
|
||||||
|
|
||||||
# Create Pydantic AI agent
|
# Create Pydantic AI agent
|
||||||
|
|
|
||||||
|
|
@ -2,7 +2,7 @@
|
||||||
|
|
||||||
name = "haiku.rag-evals"
|
name = "haiku.rag-evals"
|
||||||
description = "Benchmarking and evaluation scripts for haiku.rag"
|
description = "Benchmarking and evaluation scripts for haiku.rag"
|
||||||
version = "0.19.5"
|
version = "0.19.6"
|
||||||
authors = [{ name = "Yiorgis Gozadinos", email = "ggozadinos@gmail.com" }]
|
authors = [{ name = "Yiorgis Gozadinos", email = "ggozadinos@gmail.com" }]
|
||||||
license = { text = "MIT" }
|
license = { text = "MIT" }
|
||||||
requires-python = ">=3.12"
|
requires-python = ">=3.12"
|
||||||
|
|
|
||||||
|
|
@ -80,9 +80,10 @@ class HaikuRAGApp:
|
||||||
data = json.loads(raw) if isinstance(raw, str) else (raw or {})
|
data = json.loads(raw) if isinstance(raw, str) else (raw or {})
|
||||||
stored_version = str(data.get("version", stored_version))
|
stored_version = str(data.get("version", stored_version))
|
||||||
embeddings = data.get("embeddings", {})
|
embeddings = data.get("embeddings", {})
|
||||||
embed_provider = embeddings.get("provider")
|
embed_model_obj = embeddings.get("model", {})
|
||||||
embed_model = embeddings.get("model")
|
embed_provider = embed_model_obj.get("provider")
|
||||||
vector_dim = embeddings.get("vector_dim")
|
embed_model = embed_model_obj.get("name")
|
||||||
|
vector_dim = embed_model_obj.get("vector_dim")
|
||||||
|
|
||||||
# Get comprehensive table statistics
|
# Get comprehensive table statistics
|
||||||
from haiku.rag.store.engine import Store
|
from haiku.rag.store.engine import Store
|
||||||
|
|
|
||||||
|
|
@ -9,9 +9,11 @@ from haiku.rag.config.models import (
|
||||||
AGUIConfig,
|
AGUIConfig,
|
||||||
AppConfig,
|
AppConfig,
|
||||||
ConversionOptions,
|
ConversionOptions,
|
||||||
|
EmbeddingModelConfig,
|
||||||
EmbeddingsConfig,
|
EmbeddingsConfig,
|
||||||
LanceDBConfig,
|
LanceDBConfig,
|
||||||
LMStudioConfig,
|
LMStudioConfig,
|
||||||
|
ModelConfig,
|
||||||
MonitorConfig,
|
MonitorConfig,
|
||||||
OllamaConfig,
|
OllamaConfig,
|
||||||
ProcessingConfig,
|
ProcessingConfig,
|
||||||
|
|
@ -28,22 +30,24 @@ __all__ = [
|
||||||
"AGUIConfig",
|
"AGUIConfig",
|
||||||
"AppConfig",
|
"AppConfig",
|
||||||
"ConversionOptions",
|
"ConversionOptions",
|
||||||
"StorageConfig",
|
"EmbeddingModelConfig",
|
||||||
"MonitorConfig",
|
|
||||||
"LanceDBConfig",
|
|
||||||
"EmbeddingsConfig",
|
"EmbeddingsConfig",
|
||||||
"RerankingConfig",
|
"LanceDBConfig",
|
||||||
"QAConfig",
|
|
||||||
"ResearchConfig",
|
|
||||||
"ProcessingConfig",
|
|
||||||
"OllamaConfig",
|
|
||||||
"LMStudioConfig",
|
"LMStudioConfig",
|
||||||
"VLLMConfig",
|
"ModelConfig",
|
||||||
|
"MonitorConfig",
|
||||||
|
"OllamaConfig",
|
||||||
|
"ProcessingConfig",
|
||||||
"ProvidersConfig",
|
"ProvidersConfig",
|
||||||
|
"QAConfig",
|
||||||
|
"RerankingConfig",
|
||||||
|
"ResearchConfig",
|
||||||
|
"StorageConfig",
|
||||||
|
"VLLMConfig",
|
||||||
"find_config_file",
|
"find_config_file",
|
||||||
"load_yaml_config",
|
|
||||||
"generate_default_config",
|
"generate_default_config",
|
||||||
"get_config",
|
"get_config",
|
||||||
|
"load_yaml_config",
|
||||||
"set_config",
|
"set_config",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -25,6 +25,20 @@ class ModelConfig(BaseModel):
|
||||||
max_tokens: int | None = None
|
max_tokens: int | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class EmbeddingModelConfig(BaseModel):
|
||||||
|
"""Configuration for an embedding model.
|
||||||
|
|
||||||
|
Attributes:
|
||||||
|
provider: Model provider (ollama, openai, voyageai, vllm, lm_studio)
|
||||||
|
name: Model name/identifier
|
||||||
|
vector_dim: Vector dimensions produced by the model
|
||||||
|
"""
|
||||||
|
|
||||||
|
provider: str = "ollama"
|
||||||
|
name: str = "qwen3-embedding:4b"
|
||||||
|
vector_dim: int = 2560
|
||||||
|
|
||||||
|
|
||||||
class StorageConfig(BaseModel):
|
class StorageConfig(BaseModel):
|
||||||
data_dir: Path = Field(default_factory=get_default_data_dir)
|
data_dir: Path = Field(default_factory=get_default_data_dir)
|
||||||
vacuum_retention_seconds: int = 86400
|
vacuum_retention_seconds: int = 86400
|
||||||
|
|
@ -44,9 +58,7 @@ class LanceDBConfig(BaseModel):
|
||||||
|
|
||||||
|
|
||||||
class EmbeddingsConfig(BaseModel):
|
class EmbeddingsConfig(BaseModel):
|
||||||
provider: str = "ollama"
|
model: EmbeddingModelConfig = Field(default_factory=EmbeddingModelConfig)
|
||||||
model: str = "qwen3-embedding:4b"
|
|
||||||
vector_dim: int = 2560
|
|
||||||
|
|
||||||
|
|
||||||
class RerankingConfig(BaseModel):
|
class RerankingConfig(BaseModel):
|
||||||
|
|
|
||||||
|
|
@ -13,13 +13,12 @@ def get_embedder(config: AppConfig = Config) -> EmbedderBase:
|
||||||
Returns:
|
Returns:
|
||||||
An embedder instance configured according to the config.
|
An embedder instance configured according to the config.
|
||||||
"""
|
"""
|
||||||
|
embedding_model = config.embeddings.model
|
||||||
|
|
||||||
if config.embeddings.provider == "ollama":
|
if embedding_model.provider == "ollama":
|
||||||
return OllamaEmbedder(
|
return OllamaEmbedder(embedding_model.name, embedding_model.vector_dim, config)
|
||||||
config.embeddings.model, config.embeddings.vector_dim, config
|
|
||||||
)
|
|
||||||
|
|
||||||
if config.embeddings.provider == "voyageai":
|
if embedding_model.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:
|
||||||
|
|
@ -29,28 +28,24 @@ def get_embedder(config: AppConfig = Config) -> EmbedderBase:
|
||||||
"uv pip install haiku.rag[voyageai]"
|
"uv pip install haiku.rag[voyageai]"
|
||||||
)
|
)
|
||||||
return VoyageAIEmbedder(
|
return VoyageAIEmbedder(
|
||||||
config.embeddings.model, config.embeddings.vector_dim, config
|
embedding_model.name, embedding_model.vector_dim, config
|
||||||
)
|
)
|
||||||
|
|
||||||
if config.embeddings.provider == "openai":
|
if embedding_model.provider == "openai":
|
||||||
from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder
|
from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder
|
||||||
|
|
||||||
return OpenAIEmbedder(
|
return OpenAIEmbedder(embedding_model.name, embedding_model.vector_dim, config)
|
||||||
config.embeddings.model, config.embeddings.vector_dim, config
|
|
||||||
)
|
|
||||||
|
|
||||||
if config.embeddings.provider == "vllm":
|
if embedding_model.provider == "vllm":
|
||||||
from haiku.rag.embeddings.vllm import Embedder as VllmEmbedder
|
from haiku.rag.embeddings.vllm import Embedder as VllmEmbedder
|
||||||
|
|
||||||
return VllmEmbedder(
|
return VllmEmbedder(embedding_model.name, embedding_model.vector_dim, config)
|
||||||
config.embeddings.model, config.embeddings.vector_dim, config
|
|
||||||
)
|
|
||||||
|
|
||||||
if config.embeddings.provider == "lm_studio":
|
if embedding_model.provider == "lm_studio":
|
||||||
from haiku.rag.embeddings.lm_studio import Embedder as LMStudioEmbedder
|
from haiku.rag.embeddings.lm_studio import Embedder as LMStudioEmbedder
|
||||||
|
|
||||||
return LMStudioEmbedder(
|
return LMStudioEmbedder(
|
||||||
config.embeddings.model, config.embeddings.vector_dim, config
|
embedding_model.name, embedding_model.vector_dim, config
|
||||||
)
|
)
|
||||||
|
|
||||||
raise ValueError(f"Unsupported embedding provider: {config.embeddings.provider}")
|
raise ValueError(f"Unsupported embedding provider: {embedding_model.provider}")
|
||||||
|
|
|
||||||
|
|
@ -4,8 +4,8 @@ from haiku.rag.config import AppConfig, Config
|
||||||
|
|
||||||
|
|
||||||
class EmbedderBase:
|
class EmbedderBase:
|
||||||
_model: str = Config.embeddings.model
|
_model: str = Config.embeddings.model.name
|
||||||
_vector_dim: int = Config.embeddings.vector_dim
|
_vector_dim: int = Config.embeddings.model.vector_dim
|
||||||
_config: AppConfig = Config
|
_config: AppConfig = Config
|
||||||
|
|
||||||
def __init__(self, model: str, vector_dim: int, config: AppConfig = Config):
|
def __init__(self, model: str, vector_dim: int, config: AppConfig = Config):
|
||||||
|
|
|
||||||
|
|
@ -118,25 +118,21 @@ class SettingsRepository:
|
||||||
current_config = self.store._config.model_dump(mode="json")
|
current_config = self.store._config.model_dump(mode="json")
|
||||||
|
|
||||||
# Check if embedding provider or model has changed
|
# Check if embedding provider or model has changed
|
||||||
# Support both old flat structure and new nested structure for backward compatibility
|
# Both stored and current use nested structure: embeddings.model.{provider,name,vector_dim}
|
||||||
stored_embeddings = stored_settings.get("embeddings", {})
|
stored_embeddings = stored_settings.get("embeddings", {})
|
||||||
current_embeddings = current_config.get("embeddings", {})
|
current_embeddings = current_config.get("embeddings", {})
|
||||||
|
|
||||||
# Try nested structure first, fall back to flat for old databases
|
stored_model_obj = stored_embeddings.get("model", {})
|
||||||
stored_provider = stored_embeddings.get("provider") or stored_settings.get(
|
current_model_obj = current_embeddings.get("model", {})
|
||||||
"EMBEDDINGS_PROVIDER"
|
|
||||||
)
|
|
||||||
current_provider = current_embeddings.get("provider")
|
|
||||||
|
|
||||||
stored_model = stored_embeddings.get("model") or stored_settings.get(
|
stored_provider = stored_model_obj.get("provider")
|
||||||
"EMBEDDINGS_MODEL"
|
current_provider = current_model_obj.get("provider")
|
||||||
)
|
|
||||||
current_model = current_embeddings.get("model")
|
|
||||||
|
|
||||||
stored_vector_dim = stored_embeddings.get("vector_dim") or stored_settings.get(
|
stored_model = stored_model_obj.get("name")
|
||||||
"EMBEDDINGS_VECTOR_DIM"
|
current_model = current_model_obj.get("name")
|
||||||
)
|
|
||||||
current_vector_dim = current_embeddings.get("vector_dim")
|
stored_vector_dim = stored_model_obj.get("vector_dim")
|
||||||
|
current_vector_dim = current_model_obj.get("vector_dim")
|
||||||
|
|
||||||
# Check for incompatible changes
|
# Check for incompatible changes
|
||||||
incompatible_changes = []
|
incompatible_changes = []
|
||||||
|
|
|
||||||
|
|
@ -56,7 +56,11 @@ def run_pending_upgrades(store: Store, from_version: str, to_version: str) -> No
|
||||||
from .v0_9_3 import upgrade_fts_phrase as upgrade_0_9_3_fts # noqa: E402
|
from .v0_9_3 import upgrade_fts_phrase as upgrade_0_9_3_fts # noqa: E402
|
||||||
from .v0_9_3 import upgrade_order as upgrade_0_9_3_order # noqa: E402
|
from .v0_9_3 import upgrade_order as upgrade_0_9_3_order # noqa: E402
|
||||||
from .v0_10_1 import upgrade_add_title as upgrade_0_10_1_add_title # noqa: E402
|
from .v0_10_1 import upgrade_add_title as upgrade_0_10_1_add_title # noqa: E402
|
||||||
|
from .v0_19_6 import ( # noqa: E402
|
||||||
|
upgrade_embeddings_model_config as upgrade_0_19_6_embeddings,
|
||||||
|
)
|
||||||
|
|
||||||
upgrades.append(upgrade_0_9_3_order)
|
upgrades.append(upgrade_0_9_3_order)
|
||||||
upgrades.append(upgrade_0_9_3_fts)
|
upgrades.append(upgrade_0_9_3_fts)
|
||||||
upgrades.append(upgrade_0_10_1_add_title)
|
upgrades.append(upgrade_0_10_1_add_title)
|
||||||
|
upgrades.append(upgrade_0_19_6_embeddings)
|
||||||
|
|
|
||||||
65
haiku_rag_slim/haiku/rag/store/upgrades/v0_19_6.py
Normal file
65
haiku_rag_slim/haiku/rag/store/upgrades/v0_19_6.py
Normal file
|
|
@ -0,0 +1,65 @@
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
|
||||||
|
from haiku.rag.store.engine import SettingsRecord, Store
|
||||||
|
from haiku.rag.store.upgrades import Upgrade
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _apply_embeddings_model_config(store: Store) -> None:
|
||||||
|
"""Migrate embeddings config from flat to nested EmbeddingModelConfig structure."""
|
||||||
|
results = list(
|
||||||
|
store.settings_table.search()
|
||||||
|
.where("id = 'settings'")
|
||||||
|
.limit(1)
|
||||||
|
.to_pydantic(SettingsRecord)
|
||||||
|
)
|
||||||
|
|
||||||
|
if not results or not results[0].settings:
|
||||||
|
return
|
||||||
|
|
||||||
|
settings = json.loads(results[0].settings)
|
||||||
|
embeddings = settings.get("embeddings", {})
|
||||||
|
|
||||||
|
# Check if already migrated (model is a dict with nested structure)
|
||||||
|
if isinstance(embeddings.get("model"), dict):
|
||||||
|
return
|
||||||
|
|
||||||
|
# Migrate from flat structure to nested EmbeddingModelConfig
|
||||||
|
old_provider = embeddings.get("provider", "ollama")
|
||||||
|
old_model = embeddings.get("model", "qwen3-embedding:4b")
|
||||||
|
old_vector_dim = embeddings.get("vector_dim", 2560)
|
||||||
|
|
||||||
|
logger.warning(
|
||||||
|
"Migrating embeddings config to new nested structure: "
|
||||||
|
"embeddings.{provider,model,vector_dim} -> embeddings.model.{provider,name,vector_dim}"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Create new nested structure
|
||||||
|
settings["embeddings"] = {
|
||||||
|
"model": {
|
||||||
|
"provider": old_provider,
|
||||||
|
"name": old_model,
|
||||||
|
"vector_dim": old_vector_dim,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
store.settings_table.update(
|
||||||
|
where="id = 'settings'",
|
||||||
|
values={"settings": json.dumps(settings)},
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.warning(
|
||||||
|
"Embeddings config migrated: provider=%s, name=%s, vector_dim=%d",
|
||||||
|
old_provider,
|
||||||
|
old_model,
|
||||||
|
old_vector_dim,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
upgrade_embeddings_model_config = Upgrade(
|
||||||
|
version="0.19.6",
|
||||||
|
apply=_apply_embeddings_model_config,
|
||||||
|
description="Migrate embeddings config to nested EmbeddingModelConfig structure",
|
||||||
|
)
|
||||||
|
|
@ -392,8 +392,8 @@ async 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.model.provider == "ollama":
|
||||||
required_models.add(Config.embeddings.model)
|
required_models.add(Config.embeddings.model.name)
|
||||||
if Config.qa.model.provider == "ollama":
|
if Config.qa.model.provider == "ollama":
|
||||||
required_models.add(Config.qa.model.name)
|
required_models.add(Config.qa.model.name)
|
||||||
if Config.research.model.provider == "ollama":
|
if Config.research.model.provider == "ollama":
|
||||||
|
|
|
||||||
|
|
@ -2,7 +2,7 @@
|
||||||
|
|
||||||
name = "haiku.rag-slim"
|
name = "haiku.rag-slim"
|
||||||
description = "Opinionated agentic RAG powered by LanceDB, Pydantic AI, and Docling - Minimal dependencies"
|
description = "Opinionated agentic RAG powered by LanceDB, Pydantic AI, and Docling - Minimal dependencies"
|
||||||
version = "0.19.5"
|
version = "0.19.6"
|
||||||
authors = [{ name = "Yiorgis Gozadinos", email = "ggozadinos@gmail.com" }]
|
authors = [{ name = "Yiorgis Gozadinos", email = "ggozadinos@gmail.com" }]
|
||||||
license = { text = "MIT" }
|
license = { text = "MIT" }
|
||||||
readme = { file = "README.md", content-type = "text/markdown" }
|
readme = { file = "README.md", content-type = "text/markdown" }
|
||||||
|
|
|
||||||
|
|
@ -2,7 +2,7 @@
|
||||||
|
|
||||||
name = "haiku.rag"
|
name = "haiku.rag"
|
||||||
description = "Opinionated agentic RAG powered by LanceDB, Pydantic AI, and Docling"
|
description = "Opinionated agentic RAG powered by LanceDB, Pydantic AI, and Docling"
|
||||||
version = "0.19.5"
|
version = "0.19.6"
|
||||||
authors = [{ name = "Yiorgis Gozadinos", email = "ggozadinos@gmail.com" }]
|
authors = [{ name = "Yiorgis Gozadinos", email = "ggozadinos@gmail.com" }]
|
||||||
license = { text = "MIT" }
|
license = { text = "MIT" }
|
||||||
readme = { file = "README.md", content-type = "text/markdown" }
|
readme = { file = "README.md", content-type = "text/markdown" }
|
||||||
|
|
@ -22,7 +22,7 @@ classifiers = [
|
||||||
]
|
]
|
||||||
|
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"haiku.rag-slim[docling,voyageai,mxbai,cohere,zeroentropy,inspector]==0.19.5",
|
"haiku.rag-slim[docling,voyageai,mxbai,cohere,zeroentropy,inspector]==0.19.6",
|
||||||
]
|
]
|
||||||
|
|
||||||
[project.scripts]
|
[project.scripts]
|
||||||
|
|
|
||||||
|
|
@ -699,7 +699,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.model.vector_dim,
|
||||||
order=1,
|
order=1,
|
||||||
), # With embedding
|
), # With embedding
|
||||||
Chunk(
|
Chunk(
|
||||||
|
|
|
||||||
|
|
@ -15,16 +15,17 @@ def test_load_yaml_config(tmp_path):
|
||||||
config_file.write_text("""
|
config_file.write_text("""
|
||||||
environment: production
|
environment: production
|
||||||
embeddings:
|
embeddings:
|
||||||
provider: ollama
|
model:
|
||||||
model: test-model
|
provider: ollama
|
||||||
vector_dim: 1024
|
name: test-model
|
||||||
|
vector_dim: 1024
|
||||||
""")
|
""")
|
||||||
|
|
||||||
config = load_yaml_config(config_file)
|
config = load_yaml_config(config_file)
|
||||||
assert config["environment"] == "production"
|
assert config["environment"] == "production"
|
||||||
assert config["embeddings"]["provider"] == "ollama"
|
assert config["embeddings"]["model"]["provider"] == "ollama"
|
||||||
assert config["embeddings"]["model"] == "test-model"
|
assert config["embeddings"]["model"]["name"] == "test-model"
|
||||||
assert config["embeddings"]["vector_dim"] == 1024
|
assert config["embeddings"]["model"]["vector_dim"] == 1024
|
||||||
|
|
||||||
|
|
||||||
def test_find_config_file_cwd(tmp_path, monkeypatch):
|
def test_find_config_file_cwd(tmp_path, monkeypatch):
|
||||||
|
|
@ -179,7 +180,7 @@ def test_generate_default_config_completeness():
|
||||||
|
|
||||||
# Verify config validates successfully
|
# Verify config validates successfully
|
||||||
assert config.environment == "production"
|
assert config.environment == "production"
|
||||||
assert config.embeddings.provider == "ollama"
|
assert config.embeddings.model.provider == "ollama"
|
||||||
assert config.qa.model.provider == "ollama"
|
assert config.qa.model.provider == "ollama"
|
||||||
assert config.research.model.provider == "ollama"
|
assert config.research.model.provider == "ollama"
|
||||||
assert config.reranking.model is None
|
assert config.reranking.model is None
|
||||||
|
|
|
||||||
|
|
@ -2,6 +2,7 @@ import pytest
|
||||||
|
|
||||||
from haiku.rag.config import (
|
from haiku.rag.config import (
|
||||||
AppConfig,
|
AppConfig,
|
||||||
|
EmbeddingModelConfig,
|
||||||
EmbeddingsConfig,
|
EmbeddingsConfig,
|
||||||
LMStudioConfig,
|
LMStudioConfig,
|
||||||
OllamaConfig,
|
OllamaConfig,
|
||||||
|
|
@ -15,9 +16,9 @@ def test_embedder_uses_config_from_get_embedder():
|
||||||
"""Test that embedders use the config passed to get_embedder."""
|
"""Test that embedders use the config passed to get_embedder."""
|
||||||
custom_config = AppConfig(
|
custom_config = AppConfig(
|
||||||
embeddings=EmbeddingsConfig(
|
embeddings=EmbeddingsConfig(
|
||||||
provider="ollama",
|
model=EmbeddingModelConfig(
|
||||||
model="custom-model",
|
provider="ollama", name="custom-model", vector_dim=512
|
||||||
vector_dim=512,
|
),
|
||||||
),
|
),
|
||||||
providers=ProvidersConfig(
|
providers=ProvidersConfig(
|
||||||
ollama=OllamaConfig(base_url="http://custom-ollama:8080"),
|
ollama=OllamaConfig(base_url="http://custom-ollama:8080"),
|
||||||
|
|
@ -36,9 +37,9 @@ def test_vllm_embedder_uses_config():
|
||||||
"""Test that vllm embedder uses the config passed to get_embedder."""
|
"""Test that vllm embedder uses the config passed to get_embedder."""
|
||||||
custom_config = AppConfig(
|
custom_config = AppConfig(
|
||||||
embeddings=EmbeddingsConfig(
|
embeddings=EmbeddingsConfig(
|
||||||
provider="vllm",
|
model=EmbeddingModelConfig(
|
||||||
model="custom-vllm-model",
|
provider="vllm", name="custom-vllm-model", vector_dim=768
|
||||||
vector_dim=768,
|
),
|
||||||
),
|
),
|
||||||
providers=ProvidersConfig(
|
providers=ProvidersConfig(
|
||||||
vllm=VLLMConfig(embeddings_base_url="http://custom-vllm:9001"),
|
vllm=VLLMConfig(embeddings_base_url="http://custom-vllm:9001"),
|
||||||
|
|
@ -56,12 +57,11 @@ def test_vllm_embedder_uses_config():
|
||||||
|
|
||||||
def test_openai_embedder_uses_config():
|
def test_openai_embedder_uses_config():
|
||||||
"""Test that openai embedder uses the config passed to get_embedder."""
|
"""Test that openai embedder uses the config passed to get_embedder."""
|
||||||
|
|
||||||
custom_config = AppConfig(
|
custom_config = AppConfig(
|
||||||
embeddings=EmbeddingsConfig(
|
embeddings=EmbeddingsConfig(
|
||||||
provider="openai",
|
model=EmbeddingModelConfig(
|
||||||
model="text-embedding-3-large",
|
provider="openai", name="text-embedding-3-large", vector_dim=3072
|
||||||
vector_dim=3072,
|
),
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -76,9 +76,9 @@ def test_lm_studio_embedder_uses_config():
|
||||||
"""Test that lm_studio embedder uses the config passed to get_embedder."""
|
"""Test that lm_studio embedder uses the config passed to get_embedder."""
|
||||||
custom_config = AppConfig(
|
custom_config = AppConfig(
|
||||||
embeddings=EmbeddingsConfig(
|
embeddings=EmbeddingsConfig(
|
||||||
provider="lm_studio",
|
model=EmbeddingModelConfig(
|
||||||
model="custom-lm-studio-model",
|
provider="lm_studio", name="custom-lm-studio-model", vector_dim=1024
|
||||||
vector_dim=1024,
|
),
|
||||||
),
|
),
|
||||||
providers=ProvidersConfig(
|
providers=ProvidersConfig(
|
||||||
lm_studio=LMStudioConfig(base_url="http://custom-lmstudio:5678"),
|
lm_studio=LMStudioConfig(base_url="http://custom-lmstudio:5678"),
|
||||||
|
|
@ -101,9 +101,9 @@ def test_voyageai_embedder_uses_config():
|
||||||
"""Test that voyageai embedder uses the config passed to get_embedder."""
|
"""Test that voyageai embedder uses the config passed to get_embedder."""
|
||||||
custom_config = AppConfig(
|
custom_config = AppConfig(
|
||||||
embeddings=EmbeddingsConfig(
|
embeddings=EmbeddingsConfig(
|
||||||
provider="voyageai",
|
model=EmbeddingModelConfig(
|
||||||
model="voyage-large-2",
|
provider="voyageai", name="voyage-large-2", vector_dim=1536
|
||||||
vector_dim=1536,
|
),
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -125,7 +125,7 @@ async def test_app_info_with_vector_index(temp_db_path, capsys):
|
||||||
[
|
[
|
||||||
SettingsRecord(
|
SettingsRecord(
|
||||||
id="settings",
|
id="settings",
|
||||||
settings='{"version": "1.0.0", "embeddings": {"provider": "ollama", "model": "test", "vector_dim": 3}}',
|
settings='{"version": "1.0.0", "embeddings": {"model": {"provider": "ollama", "name": "test", "vector_dim": 3}}}',
|
||||||
)
|
)
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
|
||||||
6
uv.lock
6
uv.lock
|
|
@ -1264,7 +1264,7 @@ wheels = [
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "haiku-rag"
|
name = "haiku-rag"
|
||||||
version = "0.19.5"
|
version = "0.19.6"
|
||||||
source = { editable = "." }
|
source = { editable = "." }
|
||||||
dependencies = [
|
dependencies = [
|
||||||
{ name = "haiku-rag-slim", extra = ["cohere", "docling", "inspector", "mxbai", "voyageai", "zeroentropy"] },
|
{ name = "haiku-rag-slim", extra = ["cohere", "docling", "inspector", "mxbai", "voyageai", "zeroentropy"] },
|
||||||
|
|
@ -1312,7 +1312,7 @@ dev = [
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "haiku-rag-evals"
|
name = "haiku-rag-evals"
|
||||||
version = "0.19.5"
|
version = "0.19.6"
|
||||||
source = { editable = "evaluations" }
|
source = { editable = "evaluations" }
|
||||||
dependencies = [
|
dependencies = [
|
||||||
{ name = "datasets" },
|
{ name = "datasets" },
|
||||||
|
|
@ -1333,7 +1333,7 @@ requires-dist = [
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "haiku-rag-slim"
|
name = "haiku-rag-slim"
|
||||||
version = "0.19.5"
|
version = "0.19.6"
|
||||||
source = { editable = "haiku_rag_slim" }
|
source = { editable = "haiku_rag_slim" }
|
||||||
dependencies = [
|
dependencies = [
|
||||||
{ name = "docling-core" },
|
{ name = "docling-core" },
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue