Rename model to name under model

This commit is contained in:
Yiorgis Gozadinos 2025-11-25 12:28:43 +02:00
parent 41ea675964
commit a9701616c8
No known key found for this signature in database
17 changed files with 84 additions and 91 deletions

View file

@ -37,13 +37,13 @@ environment: production
embeddings: embeddings:
model: model:
provider: ollama provider: ollama
model: qwen3-embedding:4b name: qwen3-embedding:4b
vector_dim: 2560 vector_dim: 2560
qa: qa:
model: model:
provider: ollama provider: ollama
model: gpt-oss name: gpt-oss
enable_thinking: false enable_thinking: false
``` ```
@ -72,18 +72,18 @@ lancedb:
embeddings: embeddings:
model: model:
provider: ollama provider: ollama
model: qwen3-embedding:4b name: qwen3-embedding:4b
vector_dim: 2560 vector_dim: 2560
reranking: reranking:
model: model:
provider: "" # Empty to disable, or mxbai, cohere, zeroentropy, vllm provider: "" # Empty to disable, or mxbai, cohere, zeroentropy, vllm
model: "" name: ""
qa: qa:
model: model:
provider: ollama provider: ollama
model: gpt-oss name: gpt-oss
enable_thinking: false enable_thinking: false
max_sub_questions: 3 max_sub_questions: 3
max_iterations: 2 max_iterations: 2
@ -92,7 +92,7 @@ qa:
research: research:
model: model:
provider: "" # Empty to use qa settings provider: "" # Empty to use qa settings
model: "" name: ""
enable_thinking: true enable_thinking: true
max_iterations: 3 max_iterations: 3
confidence_threshold: 0.8 confidence_threshold: 0.8

View file

@ -15,7 +15,7 @@ Configure model behavior for `qa` and `research` workflows. These settings apply
qa: qa:
model: model:
provider: ollama provider: ollama
model: gpt-oss name: gpt-oss
temperature: 0.7 temperature: 0.7
max_tokens: 500 max_tokens: 500
``` ```
@ -74,7 +74,7 @@ If you use Ollama, you can use any pulled model that supports embeddings.
embeddings: embeddings:
model: model:
provider: ollama provider: ollama
model: mxbai-embed-large name: mxbai-embed-large
vector_dim: 1024 vector_dim: 1024
``` ```
@ -106,7 +106,7 @@ uv pip install haiku.rag-slim[voyageai]
embeddings: embeddings:
model: model:
provider: voyageai provider: voyageai
model: voyage-3.5 name: voyage-3.5
vector_dim: 1024 vector_dim: 1024
``` ```
@ -124,7 +124,7 @@ OpenAI embeddings are included in the default installation:
embeddings: embeddings:
model: model:
provider: openai provider: openai
model: text-embedding-3-small # or text-embedding-3-large name: text-embedding-3-small # or text-embedding-3-large
vector_dim: 1536 vector_dim: 1536
``` ```
@ -142,7 +142,7 @@ For high-performance local inference, you can use vLLM to serve embedding models
embeddings: embeddings:
model: model:
provider: vllm provider: vllm
model: mixedbread-ai/mxbai-embed-large-v1 name: mixedbread-ai/mxbai-embed-large-v1
vector_dim: 512 vector_dim: 512
providers: providers:
@ -162,7 +162,7 @@ Configure which LLM provider to use for question answering. Any provider and mod
qa: qa:
model: model:
provider: ollama provider: ollama
model: gpt-oss name: gpt-oss
``` ```
The Ollama base URL can be configured via the `OLLAMA_BASE_URL` environment variable, config file, or defaults to `http://localhost:11434`: The Ollama base URL can be configured via the `OLLAMA_BASE_URL` environment variable, config file, or defaults to `http://localhost:11434`:
@ -187,7 +187,7 @@ OpenAI QA is included in the default installation:
qa: qa:
model: model:
provider: openai provider: openai
model: gpt-4o-mini # or gpt-4, gpt-3.5-turbo, etc. name: gpt-4o-mini # or gpt-4, gpt-3.5-turbo, etc.
``` ```
Set your API key via environment variable: Set your API key via environment variable:
@ -204,7 +204,7 @@ Anthropic QA is included in the default installation:
qa: qa:
model: model:
provider: anthropic provider: anthropic
model: claude-3-5-haiku-20241022 # or claude-3-5-sonnet-20241022, etc. name: claude-3-5-haiku-20241022 # or claude-3-5-sonnet-20241022, etc.
``` ```
Set your API key via environment variable: Set your API key via environment variable:
@ -221,7 +221,7 @@ For high-performance local inference:
qa: qa:
model: model:
provider: vllm provider: vllm
model: Qwen/Qwen3-4B # Any model with tool support in vLLM name: Qwen/Qwen3-4B # Any model with tool support in vLLM
providers: providers:
vllm: vllm:
@ -239,19 +239,19 @@ Any provider supported by Pydantic AI can be used. Examples:
qa: qa:
model: model:
provider: gemini provider: gemini
model: gemini-1.5-flash name: gemini-1.5-flash
# Groq # Groq
qa: qa:
model: model:
provider: groq provider: groq
model: llama-3.3-70b-versatile name: llama-3.3-70b-versatile
# Mistral # Mistral
qa: qa:
model: model:
provider: mistral provider: mistral
model: mistral-small-latest name: 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.
@ -276,7 +276,7 @@ Then configure:
reranking: reranking:
model: model:
provider: mxbai provider: mxbai
model: mixedbread-ai/mxbai-rerank-base-v2 name: mixedbread-ai/mxbai-rerank-base-v2
``` ```
### Cohere ### Cohere
@ -293,7 +293,7 @@ Then configure:
reranking: reranking:
model: model:
provider: cohere provider: cohere
model: rerank-v3.5 name: rerank-v3.5
``` ```
Set your API key via environment variable: Set your API key via environment variable:
@ -316,7 +316,7 @@ Then configure:
reranking: reranking:
model: model:
provider: zeroentropy provider: zeroentropy
model: zerank-1 # Currently the only available model name: zerank-1 # Currently the only available model
``` ```
Set your API key via environment variable: Set your API key via environment variable:
@ -333,7 +333,7 @@ For high-performance local reranking using dedicated reranking models:
reranking: reranking:
model: model:
provider: vllm provider: vllm
model: mixedbread-ai/mxbai-rerank-base-v2 name: mixedbread-ai/mxbai-rerank-base-v2
providers: providers:
vllm: vllm:

View file

@ -8,7 +8,7 @@ Configure the QA workflow:
qa: qa:
model: model:
provider: ollama provider: ollama
model: gpt-oss name: gpt-oss
enable_thinking: false enable_thinking: false
max_sub_questions: 3 # Maximum sub-questions for deep QA max_sub_questions: 3 # Maximum sub-questions for deep QA
max_iterations: 2 # Maximum search iterations per sub-question max_iterations: 2 # Maximum search iterations per sub-question
@ -30,7 +30,7 @@ Configure the multi-agent research workflow:
research: research:
model: model:
provider: "" # Empty to use qa settings provider: "" # Empty to use qa settings
model: "" # Empty to use qa model name: "" # Empty to use qa model
enable_thinking: true enable_thinking: true
max_iterations: 3 max_iterations: 3
confidence_threshold: 0.8 confidence_threshold: 0.8

View file

@ -44,18 +44,16 @@ def build_experiment_metadata(
"dataset": dataset_key, "dataset": dataset_key,
"test_cases": test_cases, "test_cases": test_cases,
"embedder_provider": config.embeddings.model.provider, "embedder_provider": config.embeddings.model.provider,
"embedder_model": config.embeddings.model.model, "embedder_model": config.embeddings.model.name,
"embedder_dim": config.embeddings.vector_dim, "embedder_dim": config.embeddings.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
if config.reranking.model if config.reranking.model
else None, else None,
"rerank_model": config.reranking.model.model "rerank_model": config.reranking.model.name if config.reranking.model else None,
if config.reranking.model
else None,
"qa_provider": config.qa.model.provider, "qa_provider": config.qa.model.provider,
"qa_model": config.qa.model.model, "qa_model": config.qa.model.name,
"judge_provider": "ollama", "judge_provider": "ollama",
"judge_model": judge_model, "judge_model": judge_model,
} }

View file

@ -47,7 +47,7 @@ if not db_path.exists():
logger.info(f"Initializing research assistant with database: {db_path}") logger.info(f"Initializing research assistant with database: {db_path}")
logger.info( logger.info(
f"Research Provider: {Config.research.model.provider}, Model: {Config.research.model.model}" f"Research Provider: {Config.research.model.provider}, Model: {Config.research.model.name}"
) )
# Store client reference for proper lifecycle management # Store client reference for proper lifecycle management
@ -154,7 +154,7 @@ async def health_check(_: Request) -> JSONResponse:
"status": "healthy", "status": "healthy",
"agent_model": str(agent.model), "agent_model": str(agent.model),
"research_provider": Config.research.model.provider, "research_provider": Config.research.model.provider,
"research_model": Config.research.model.model, "research_model": Config.research.model.name,
"db_path": str(db_path), "db_path": str(db_path),
"db_exists": db_path.exists(), "db_exists": db_path.exists(),
} }

View file

@ -17,18 +17,21 @@ providers:
base_url: http://host.docker.internal:11434 base_url: http://host.docker.internal:11434
research: research:
provider: ollama model:
model: gpt-oss:latest provider: ollama
name: gpt-oss:latest
max_iterations: 3 max_iterations: 3
confidence_threshold: 0.8 confidence_threshold: 0.8
max_concurrency: 1 max_concurrency: 1
# For OpenAI: # For OpenAI:
# research: # research:
# provider: openai # model:
# model: gpt-4o-mini # provider: openai
# name: gpt-4o-mini
# For Anthropic: # For Anthropic:
# research: # research:
# provider: anthropic # model:
# model: claude-3-5-haiku-20241022 # provider: anthropic
# name: claude-3-5-haiku-20241022

View file

@ -11,14 +11,14 @@ class ModelConfig(BaseModel):
Attributes: Attributes:
provider: Model provider (ollama, openai, anthropic, etc.) provider: Model provider (ollama, openai, anthropic, etc.)
model: Model name/identifier name: Model name/identifier
enable_thinking: Control reasoning behavior (true/false/None for default) enable_thinking: Control reasoning behavior (true/false/None for default)
temperature: Sampling temperature (0.0 to 1.0+) temperature: Sampling temperature (0.0 to 1.0+)
max_tokens: Maximum tokens to generate max_tokens: Maximum tokens to generate
""" """
provider: str = "ollama" provider: str = "ollama"
model: str = "gpt-oss" name: str = "gpt-oss"
enable_thinking: bool | None = None enable_thinking: bool | None = None
temperature: float | None = None temperature: float | None = None
@ -47,7 +47,7 @@ class EmbeddingsConfig(BaseModel):
model: ModelConfig = Field( model: ModelConfig = Field(
default_factory=lambda: ModelConfig( default_factory=lambda: ModelConfig(
provider="ollama", provider="ollama",
model="qwen3-embedding:4b", name="qwen3-embedding:4b",
) )
) )
vector_dim: int = 2560 vector_dim: int = 2560
@ -61,7 +61,7 @@ class QAConfig(BaseModel):
model: ModelConfig = Field( model: ModelConfig = Field(
default_factory=lambda: ModelConfig( default_factory=lambda: ModelConfig(
provider="ollama", provider="ollama",
model="gpt-oss", name="gpt-oss",
enable_thinking=False, enable_thinking=False,
) )
) )
@ -74,7 +74,7 @@ class ResearchConfig(BaseModel):
model: ModelConfig = Field( model: ModelConfig = Field(
default_factory=lambda: ModelConfig( default_factory=lambda: ModelConfig(
provider="ollama", provider="ollama",
model="gpt-oss", name="gpt-oss",
enable_thinking=True, enable_thinking=True,
) )
) )

View file

@ -16,7 +16,7 @@ def get_embedder(config: AppConfig = Config) -> EmbedderBase:
if config.embeddings.model.provider == "ollama": if config.embeddings.model.provider == "ollama":
return OllamaEmbedder( return OllamaEmbedder(
config.embeddings.model.model, config.embeddings.vector_dim, config config.embeddings.model.name, config.embeddings.vector_dim, config
) )
if config.embeddings.model.provider == "voyageai": if config.embeddings.model.provider == "voyageai":
@ -29,21 +29,21 @@ 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.model, config.embeddings.vector_dim, config config.embeddings.model.name, config.embeddings.vector_dim, config
) )
if config.embeddings.model.provider == "openai": if config.embeddings.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(
config.embeddings.model.model, config.embeddings.vector_dim, config config.embeddings.model.name, config.embeddings.vector_dim, config
) )
if config.embeddings.model.provider == "vllm": if config.embeddings.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(
config.embeddings.model.model, config.embeddings.vector_dim, config config.embeddings.model.name, config.embeddings.vector_dim, config
) )
raise ValueError( raise ValueError(

View file

@ -4,7 +4,7 @@ from haiku.rag.config import AppConfig, Config
class EmbedderBase: class EmbedderBase:
_model: str = Config.embeddings.model.model _model: str = Config.embeddings.model.name
_vector_dim: int = Config.embeddings.vector_dim _vector_dim: int = Config.embeddings.vector_dim
_config: AppConfig = Config _config: AppConfig = Config

View file

@ -45,7 +45,7 @@ def get_reranker(config: AppConfig = Config) -> RerankerBase | None:
try: try:
from haiku.rag.reranking.vllm import VLLMReranker from haiku.rag.reranking.vllm import VLLMReranker
reranker = VLLMReranker(config.reranking.model.model) reranker = VLLMReranker(config.reranking.model.name)
except ImportError: except ImportError:
reranker = None reranker = None
@ -54,7 +54,7 @@ def get_reranker(config: AppConfig = Config) -> RerankerBase | None:
from haiku.rag.reranking.zeroentropy import ZeroEntropyReranker from haiku.rag.reranking.zeroentropy import ZeroEntropyReranker
# Use configured model or default to zerank-1 # Use configured model or default to zerank-1
model = config.reranking.model.model or "zerank-1" model = config.reranking.model.name or "zerank-1"
reranker = ZeroEntropyReranker(model) reranker = ZeroEntropyReranker(model)
except ImportError: except ImportError:
reranker = None reranker = None

View file

@ -3,9 +3,7 @@ from haiku.rag.store.models.chunk import Chunk
class RerankerBase: class RerankerBase:
_model: str | None = ( _model: str | None = Config.reranking.model.name if Config.reranking.model else None
Config.reranking.model.model if Config.reranking.model else None
)
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

View file

@ -8,7 +8,7 @@ from haiku.rag.store.models.chunk import Chunk
class MxBAIReranker(RerankerBase): class MxBAIReranker(RerankerBase):
def __init__(self): def __init__(self):
model_name = ( model_name = (
Config.reranking.model.model Config.reranking.model.name
if Config.reranking.model if Config.reranking.model
else "mxbai-rerank-base-v2" else "mxbai-rerank-base-v2"
) )

View file

@ -65,7 +65,7 @@ def get_model(
app_config = Config app_config = Config
provider = model_config.provider provider = model_config.provider
model = model_config.model model = model_config.name
if provider == "ollama": if provider == "ollama":
model_settings = None model_settings = None
@ -366,13 +366,13 @@ 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.model.provider == "ollama": if Config.embeddings.model.provider == "ollama":
required_models.add(Config.embeddings.model.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.model) required_models.add(Config.qa.model.name)
if Config.research.model.provider == "ollama": if Config.research.model.provider == "ollama":
required_models.add(Config.research.model.model) required_models.add(Config.research.model.name)
if Config.reranking.model and Config.reranking.model.provider == "ollama": if Config.reranking.model and Config.reranking.model.provider == "ollama":
required_models.add(Config.reranking.model.model) required_models.add(Config.reranking.model.name)
if not required_models: if not required_models:
return return

View file

@ -17,7 +17,7 @@ def test_embedder_uses_config_from_get_embedder():
embeddings=EmbeddingsConfig( embeddings=EmbeddingsConfig(
model=ModelConfig( model=ModelConfig(
provider="ollama", provider="ollama",
model="custom-model", name="custom-model",
), ),
vector_dim=512, vector_dim=512,
), ),
@ -42,7 +42,7 @@ def test_vllm_embedder_uses_config():
embeddings=EmbeddingsConfig( embeddings=EmbeddingsConfig(
model=ModelConfig( model=ModelConfig(
provider="vllm", provider="vllm",
model="custom-vllm-model", name="custom-vllm-model",
), ),
vector_dim=768, vector_dim=768,
), ),
@ -68,7 +68,7 @@ def test_openai_embedder_uses_config():
embeddings=EmbeddingsConfig( embeddings=EmbeddingsConfig(
model=ModelConfig( model=ModelConfig(
provider="openai", provider="openai",
model="text-embedding-3-large", name="text-embedding-3-large",
), ),
vector_dim=3072, vector_dim=3072,
), ),
@ -90,7 +90,7 @@ def test_voyageai_embedder_uses_config():
embeddings=EmbeddingsConfig( embeddings=EmbeddingsConfig(
model=ModelConfig( model=ModelConfig(
provider="voyageai", provider="voyageai",
model="voyage-large-2", name="voyage-large-2",
), ),
vector_dim=1536, vector_dim=1536,
), ),

View file

@ -19,7 +19,7 @@ async def test_qa_ollama(qa_corpus: Dataset, temp_db_path):
"""Test Ollama QA with LLM judge.""" """Test Ollama QA with LLM judge."""
client = HaikuRAG(temp_db_path) client = HaikuRAG(temp_db_path)
qa = QuestionAnswerAgent( qa = QuestionAnswerAgent(
client, ModelConfig(provider="ollama", model="gpt-oss", enable_thinking=False) client, ModelConfig(provider="ollama", name="gpt-oss", enable_thinking=False)
) )
llm_judge = LLMJudge() llm_judge = LLMJudge()
@ -44,9 +44,7 @@ async def test_qa_ollama(qa_corpus: Dataset, temp_db_path):
async def test_qa_openai(qa_corpus: Dataset, temp_db_path): async def test_qa_openai(qa_corpus: Dataset, temp_db_path):
"""Test OpenAI QA with LLM judge.""" """Test OpenAI QA with LLM judge."""
client = HaikuRAG(temp_db_path) client = HaikuRAG(temp_db_path)
qa = QuestionAnswerAgent( qa = QuestionAnswerAgent(client, ModelConfig(provider="openai", name="gpt-4o-mini"))
client, ModelConfig(provider="openai", model="gpt-4o-mini")
)
llm_judge = LLMJudge() llm_judge = LLMJudge()
doc = qa_corpus[1] doc = qa_corpus[1]
@ -71,7 +69,7 @@ async def test_qa_anthropic(qa_corpus: Dataset, temp_db_path):
"""Test Anthropic QA with LLM judge.""" """Test Anthropic QA with LLM judge."""
client = HaikuRAG(temp_db_path) client = HaikuRAG(temp_db_path)
qa = QuestionAnswerAgent( qa = QuestionAnswerAgent(
client, ModelConfig(provider="anthropic", model="claude-3-5-haiku-20241022") client, ModelConfig(provider="anthropic", name="claude-3-5-haiku-20241022")
) )
llm_judge = LLMJudge() llm_judge = LLMJudge()
@ -96,9 +94,7 @@ async def test_qa_anthropic(qa_corpus: Dataset, temp_db_path):
async def test_qa_vllm(qa_corpus: Dataset, temp_db_path): async def test_qa_vllm(qa_corpus: Dataset, temp_db_path):
"""Test vLLM QA with LLM judge.""" """Test vLLM QA with LLM judge."""
client = HaikuRAG(temp_db_path) client = HaikuRAG(temp_db_path)
qa = QuestionAnswerAgent( qa = QuestionAnswerAgent(client, ModelConfig(provider="vllm", name="Qwen/Qwen3-4B"))
client, ModelConfig(provider="vllm", model="Qwen/Qwen3-4B")
)
llm_judge = LLMJudge() llm_judge = LLMJudge()
doc = qa_corpus[1] doc = qa_corpus[1]

View file

@ -44,7 +44,7 @@ async def test_mxbai_reranker():
from haiku.rag.reranking.mxbai import MxBAIReranker from haiku.rag.reranking.mxbai import MxBAIReranker
Config.reranking.model = ModelConfig( Config.reranking.model = ModelConfig(
provider="mxbai", model="mixedbread-ai/mxbai-rerank-base-v2" provider="mxbai", name="mixedbread-ai/mxbai-rerank-base-v2"
) )
reranker = MxBAIReranker() reranker = MxBAIReranker()
reranked = await reranker.rerank( reranked = await reranker.rerank(

View file

@ -136,16 +136,14 @@ Emoji test: 🚀 ✅ 📝"""
def test_get_model_ollama(): def test_get_model_ollama():
"""Test get_model returns OpenAIChatModel for Ollama.""" """Test get_model returns OpenAIChatModel for Ollama."""
model_config = ModelConfig(provider="ollama", model="llama3") model_config = ModelConfig(provider="ollama", name="llama3")
result = get_model(model_config) result = get_model(model_config)
assert isinstance(result, OpenAIChatModel) assert isinstance(result, OpenAIChatModel)
def test_get_model_ollama_with_thinking(): def test_get_model_ollama_without_thinking():
"""Test get_model configures thinking for gpt-oss on Ollama.""" """Test get_model configures thinking for gpt-oss on Ollama."""
model_config = ModelConfig( model_config = ModelConfig(provider="ollama", name="gpt-oss", enable_thinking=False)
provider="ollama", model="gpt-oss", enable_thinking=False
)
result = get_model(model_config) result = get_model(model_config)
assert isinstance(result, OpenAIChatModel) assert isinstance(result, OpenAIChatModel)
@ -153,7 +151,7 @@ def test_get_model_ollama_with_thinking():
def test_get_model_ollama_with_settings(): def test_get_model_ollama_with_settings():
"""Test get_model applies temperature and max_tokens for Ollama.""" """Test get_model applies temperature and max_tokens for Ollama."""
model_config = ModelConfig( model_config = ModelConfig(
provider="ollama", model="llama3", temperature=0.5, max_tokens=100 provider="ollama", name="llama3", temperature=0.5, max_tokens=100
) )
result = get_model(model_config) result = get_model(model_config)
assert isinstance(result, OpenAIChatModel) assert isinstance(result, OpenAIChatModel)
@ -161,14 +159,14 @@ def test_get_model_ollama_with_settings():
def test_get_model_openai(): def test_get_model_openai():
"""Test get_model returns OpenAIChatModel for OpenAI.""" """Test get_model returns OpenAIChatModel for OpenAI."""
model_config = ModelConfig(provider="openai", model="gpt-4o") model_config = ModelConfig(provider="openai", name="gpt-4o")
result = get_model(model_config) result = get_model(model_config)
assert isinstance(result, OpenAIChatModel) assert isinstance(result, OpenAIChatModel)
def test_get_model_openai_with_thinking(): def test_get_model_openai_with_thinking():
"""Test get_model configures thinking for OpenAI reasoning models.""" """Test get_model configures thinking for OpenAI reasoning models."""
model_config = ModelConfig(provider="openai", model="o1", enable_thinking=True) model_config = ModelConfig(provider="openai", name="o1", enable_thinking=True)
result = get_model(model_config) result = get_model(model_config)
assert isinstance(result, OpenAIChatModel) assert isinstance(result, OpenAIChatModel)
@ -178,7 +176,7 @@ def test_get_model_anthropic():
"""Test get_model returns AnthropicModel for Anthropic.""" """Test get_model returns AnthropicModel for Anthropic."""
from pydantic_ai.models.anthropic import AnthropicModel from pydantic_ai.models.anthropic import AnthropicModel
model_config = ModelConfig(provider="anthropic", model="claude-3-5-sonnet-20241022") model_config = ModelConfig(provider="anthropic", name="claude-3-5-sonnet-20241022")
result = get_model(model_config) result = get_model(model_config)
assert isinstance(result, AnthropicModel) assert isinstance(result, AnthropicModel)
@ -190,7 +188,7 @@ def test_get_model_anthropic_with_thinking():
model_config = ModelConfig( model_config = ModelConfig(
provider="anthropic", provider="anthropic",
model="claude-3-5-sonnet-20241022", name="claude-3-5-sonnet-20241022",
enable_thinking=True, enable_thinking=True,
) )
result = get_model(model_config) result = get_model(model_config)
@ -202,7 +200,7 @@ def test_get_model_gemini():
"""Test get_model returns GoogleModel for Gemini.""" """Test get_model returns GoogleModel for Gemini."""
from pydantic_ai.models.google import GoogleModel from pydantic_ai.models.google import GoogleModel
model_config = ModelConfig(provider="gemini", model="gemini-2.0-flash-exp") model_config = ModelConfig(provider="gemini", name="gemini-2.0-flash-exp")
result = get_model(model_config) result = get_model(model_config)
assert isinstance(result, GoogleModel) assert isinstance(result, GoogleModel)
@ -213,7 +211,7 @@ def test_get_model_gemini_with_thinking():
from pydantic_ai.models.google import GoogleModel from pydantic_ai.models.google import GoogleModel
model_config = ModelConfig( model_config = ModelConfig(
provider="gemini", model="gemini-2.0-flash-thinking-exp", enable_thinking=True provider="gemini", name="gemini-2.0-flash-thinking-exp", enable_thinking=True
) )
result = get_model(model_config) result = get_model(model_config)
assert isinstance(result, GoogleModel) assert isinstance(result, GoogleModel)
@ -224,7 +222,7 @@ def test_get_model_groq():
"""Test get_model returns GroqModel for Groq.""" """Test get_model returns GroqModel for Groq."""
from pydantic_ai.models.groq import GroqModel from pydantic_ai.models.groq import GroqModel
model_config = ModelConfig(provider="groq", model="llama-3.3-70b-versatile") model_config = ModelConfig(provider="groq", name="llama-3.3-70b-versatile")
result = get_model(model_config) result = get_model(model_config)
assert isinstance(result, GroqModel) assert isinstance(result, GroqModel)
@ -235,7 +233,7 @@ def test_get_model_groq_with_thinking():
from pydantic_ai.models.groq import GroqModel from pydantic_ai.models.groq import GroqModel
model_config = ModelConfig( model_config = ModelConfig(
provider="groq", model="llama-3.3-70b-versatile", enable_thinking=False provider="groq", name="llama-3.3-70b-versatile", enable_thinking=False
) )
result = get_model(model_config) result = get_model(model_config)
assert isinstance(result, GroqModel) assert isinstance(result, GroqModel)
@ -247,7 +245,7 @@ def test_get_model_bedrock():
from pydantic_ai.models.bedrock import BedrockConverseModel from pydantic_ai.models.bedrock import BedrockConverseModel
model_config = ModelConfig( model_config = ModelConfig(
provider="bedrock", model="anthropic.claude-3-5-sonnet-20241022-v2:0" provider="bedrock", name="anthropic.claude-3-5-sonnet-20241022-v2:0"
) )
result = get_model(model_config) result = get_model(model_config)
assert isinstance(result, BedrockConverseModel) assert isinstance(result, BedrockConverseModel)
@ -260,7 +258,7 @@ def test_get_model_bedrock_with_thinking():
model_config = ModelConfig( model_config = ModelConfig(
provider="bedrock", provider="bedrock",
model="anthropic.claude-3-5-sonnet-20241022-v2:0", name="anthropic.claude-3-5-sonnet-20241022-v2:0",
enable_thinking=True, enable_thinking=True,
) )
result = get_model(model_config) result = get_model(model_config)
@ -269,21 +267,21 @@ def test_get_model_bedrock_with_thinking():
def test_get_model_vllm(): def test_get_model_vllm():
"""Test get_model returns OpenAIChatModel for vLLM.""" """Test get_model returns OpenAIChatModel for vLLM."""
model_config = ModelConfig(provider="vllm", model="Qwen/Qwen3-4B") model_config = ModelConfig(provider="vllm", name="Qwen/Qwen3-4B")
result = get_model(model_config) result = get_model(model_config)
assert isinstance(result, OpenAIChatModel) assert isinstance(result, OpenAIChatModel)
def test_get_model_vllm_with_thinking(): def test_get_model_vllm_with_thinking():
"""Test get_model configures thinking for gpt-oss on vLLM.""" """Test get_model configures thinking for gpt-oss on vLLM."""
model_config = ModelConfig(provider="vllm", model="gpt-oss", enable_thinking=False) model_config = ModelConfig(provider="vllm", name="gpt-oss", enable_thinking=False)
result = get_model(model_config) result = get_model(model_config)
assert isinstance(result, OpenAIChatModel) assert isinstance(result, OpenAIChatModel)
def test_get_model_unknown_provider(): def test_get_model_unknown_provider():
"""Test get_model returns string format for unknown providers.""" """Test get_model returns string format for unknown providers."""
model_config = ModelConfig(provider="mistral", model="mistral-large-latest") model_config = ModelConfig(provider="mistral", name="mistral-large-latest")
result = get_model(model_config) result = get_model(model_config)
assert isinstance(result, str) assert isinstance(result, str)
assert result == "mistral:mistral-large-latest" assert result == "mistral:mistral-large-latest"
@ -293,7 +291,7 @@ def test_get_model_with_all_settings():
"""Test get_model applies all settings together.""" """Test get_model applies all settings together."""
model_config = ModelConfig( model_config = ModelConfig(
provider="openai", provider="openai",
model="gpt-4o", name="gpt-4o",
enable_thinking=False, enable_thinking=False,
temperature=0.7, temperature=0.7,
max_tokens=500, max_tokens=500,