Rename model to name under model
This commit is contained in:
parent
41ea675964
commit
a9701616c8
17 changed files with 84 additions and 91 deletions
|
|
@ -37,13 +37,13 @@ environment: production
|
|||
embeddings:
|
||||
model:
|
||||
provider: ollama
|
||||
model: qwen3-embedding:4b
|
||||
name: qwen3-embedding:4b
|
||||
vector_dim: 2560
|
||||
|
||||
qa:
|
||||
model:
|
||||
provider: ollama
|
||||
model: gpt-oss
|
||||
name: gpt-oss
|
||||
enable_thinking: false
|
||||
```
|
||||
|
||||
|
|
@ -72,18 +72,18 @@ lancedb:
|
|||
embeddings:
|
||||
model:
|
||||
provider: ollama
|
||||
model: qwen3-embedding:4b
|
||||
name: qwen3-embedding:4b
|
||||
vector_dim: 2560
|
||||
|
||||
reranking:
|
||||
model:
|
||||
provider: "" # Empty to disable, or mxbai, cohere, zeroentropy, vllm
|
||||
model: ""
|
||||
name: ""
|
||||
|
||||
qa:
|
||||
model:
|
||||
provider: ollama
|
||||
model: gpt-oss
|
||||
name: gpt-oss
|
||||
enable_thinking: false
|
||||
max_sub_questions: 3
|
||||
max_iterations: 2
|
||||
|
|
@ -92,7 +92,7 @@ qa:
|
|||
research:
|
||||
model:
|
||||
provider: "" # Empty to use qa settings
|
||||
model: ""
|
||||
name: ""
|
||||
enable_thinking: true
|
||||
max_iterations: 3
|
||||
confidence_threshold: 0.8
|
||||
|
|
|
|||
|
|
@ -15,7 +15,7 @@ Configure model behavior for `qa` and `research` workflows. These settings apply
|
|||
qa:
|
||||
model:
|
||||
provider: ollama
|
||||
model: gpt-oss
|
||||
name: gpt-oss
|
||||
temperature: 0.7
|
||||
max_tokens: 500
|
||||
```
|
||||
|
|
@ -74,7 +74,7 @@ If you use Ollama, you can use any pulled model that supports embeddings.
|
|||
embeddings:
|
||||
model:
|
||||
provider: ollama
|
||||
model: mxbai-embed-large
|
||||
name: mxbai-embed-large
|
||||
vector_dim: 1024
|
||||
```
|
||||
|
||||
|
|
@ -106,7 +106,7 @@ uv pip install haiku.rag-slim[voyageai]
|
|||
embeddings:
|
||||
model:
|
||||
provider: voyageai
|
||||
model: voyage-3.5
|
||||
name: voyage-3.5
|
||||
vector_dim: 1024
|
||||
```
|
||||
|
||||
|
|
@ -124,7 +124,7 @@ OpenAI embeddings are included in the default installation:
|
|||
embeddings:
|
||||
model:
|
||||
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
|
||||
```
|
||||
|
||||
|
|
@ -142,7 +142,7 @@ For high-performance local inference, you can use vLLM to serve embedding models
|
|||
embeddings:
|
||||
model:
|
||||
provider: vllm
|
||||
model: mixedbread-ai/mxbai-embed-large-v1
|
||||
name: mixedbread-ai/mxbai-embed-large-v1
|
||||
vector_dim: 512
|
||||
|
||||
providers:
|
||||
|
|
@ -162,7 +162,7 @@ Configure which LLM provider to use for question answering. Any provider and mod
|
|||
qa:
|
||||
model:
|
||||
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`:
|
||||
|
|
@ -187,7 +187,7 @@ OpenAI QA is included in the default installation:
|
|||
qa:
|
||||
model:
|
||||
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:
|
||||
|
|
@ -204,7 +204,7 @@ Anthropic QA is included in the default installation:
|
|||
qa:
|
||||
model:
|
||||
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:
|
||||
|
|
@ -221,7 +221,7 @@ For high-performance local inference:
|
|||
qa:
|
||||
model:
|
||||
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:
|
||||
vllm:
|
||||
|
|
@ -239,19 +239,19 @@ Any provider supported by Pydantic AI can be used. Examples:
|
|||
qa:
|
||||
model:
|
||||
provider: gemini
|
||||
model: gemini-1.5-flash
|
||||
name: gemini-1.5-flash
|
||||
|
||||
# Groq
|
||||
qa:
|
||||
model:
|
||||
provider: groq
|
||||
model: llama-3.3-70b-versatile
|
||||
name: llama-3.3-70b-versatile
|
||||
|
||||
# Mistral
|
||||
qa:
|
||||
model:
|
||||
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.
|
||||
|
|
@ -276,7 +276,7 @@ Then configure:
|
|||
reranking:
|
||||
model:
|
||||
provider: mxbai
|
||||
model: mixedbread-ai/mxbai-rerank-base-v2
|
||||
name: mixedbread-ai/mxbai-rerank-base-v2
|
||||
```
|
||||
|
||||
### Cohere
|
||||
|
|
@ -293,7 +293,7 @@ Then configure:
|
|||
reranking:
|
||||
model:
|
||||
provider: cohere
|
||||
model: rerank-v3.5
|
||||
name: rerank-v3.5
|
||||
```
|
||||
|
||||
Set your API key via environment variable:
|
||||
|
|
@ -316,7 +316,7 @@ Then configure:
|
|||
reranking:
|
||||
model:
|
||||
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:
|
||||
|
|
@ -333,7 +333,7 @@ For high-performance local reranking using dedicated reranking models:
|
|||
reranking:
|
||||
model:
|
||||
provider: vllm
|
||||
model: mixedbread-ai/mxbai-rerank-base-v2
|
||||
name: mixedbread-ai/mxbai-rerank-base-v2
|
||||
|
||||
providers:
|
||||
vllm:
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ Configure the QA workflow:
|
|||
qa:
|
||||
model:
|
||||
provider: ollama
|
||||
model: gpt-oss
|
||||
name: gpt-oss
|
||||
enable_thinking: false
|
||||
max_sub_questions: 3 # Maximum sub-questions for deep QA
|
||||
max_iterations: 2 # Maximum search iterations per sub-question
|
||||
|
|
@ -30,7 +30,7 @@ Configure the multi-agent research workflow:
|
|||
research:
|
||||
model:
|
||||
provider: "" # Empty to use qa settings
|
||||
model: "" # Empty to use qa model
|
||||
name: "" # Empty to use qa model
|
||||
enable_thinking: true
|
||||
max_iterations: 3
|
||||
confidence_threshold: 0.8
|
||||
|
|
|
|||
|
|
@ -44,18 +44,16 @@ def build_experiment_metadata(
|
|||
"dataset": dataset_key,
|
||||
"test_cases": test_cases,
|
||||
"embedder_provider": config.embeddings.model.provider,
|
||||
"embedder_model": config.embeddings.model.model,
|
||||
"embedder_model": config.embeddings.model.name,
|
||||
"embedder_dim": config.embeddings.vector_dim,
|
||||
"chunk_size": config.processing.chunk_size,
|
||||
"context_chunk_radius": config.processing.context_chunk_radius,
|
||||
"rerank_provider": config.reranking.model.provider
|
||||
if config.reranking.model
|
||||
else None,
|
||||
"rerank_model": config.reranking.model.model
|
||||
if config.reranking.model
|
||||
else None,
|
||||
"rerank_model": config.reranking.model.name if config.reranking.model else None,
|
||||
"qa_provider": config.qa.model.provider,
|
||||
"qa_model": config.qa.model.model,
|
||||
"qa_model": config.qa.model.name,
|
||||
"judge_provider": "ollama",
|
||||
"judge_model": judge_model,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -47,7 +47,7 @@ if not db_path.exists():
|
|||
|
||||
logger.info(f"Initializing research assistant with database: {db_path}")
|
||||
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
|
||||
|
|
@ -154,7 +154,7 @@ async def health_check(_: Request) -> JSONResponse:
|
|||
"status": "healthy",
|
||||
"agent_model": str(agent.model),
|
||||
"research_provider": Config.research.model.provider,
|
||||
"research_model": Config.research.model.model,
|
||||
"research_model": Config.research.model.name,
|
||||
"db_path": str(db_path),
|
||||
"db_exists": db_path.exists(),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -17,18 +17,21 @@ providers:
|
|||
base_url: http://host.docker.internal:11434
|
||||
|
||||
research:
|
||||
provider: ollama
|
||||
model: gpt-oss:latest
|
||||
model:
|
||||
provider: ollama
|
||||
name: gpt-oss:latest
|
||||
max_iterations: 3
|
||||
confidence_threshold: 0.8
|
||||
max_concurrency: 1
|
||||
|
||||
# For OpenAI:
|
||||
# research:
|
||||
# provider: openai
|
||||
# model: gpt-4o-mini
|
||||
# model:
|
||||
# provider: openai
|
||||
# name: gpt-4o-mini
|
||||
|
||||
# For Anthropic:
|
||||
# research:
|
||||
# provider: anthropic
|
||||
# model: claude-3-5-haiku-20241022
|
||||
# model:
|
||||
# provider: anthropic
|
||||
# name: claude-3-5-haiku-20241022
|
||||
|
|
|
|||
|
|
@ -11,14 +11,14 @@ class ModelConfig(BaseModel):
|
|||
|
||||
Attributes:
|
||||
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)
|
||||
temperature: Sampling temperature (0.0 to 1.0+)
|
||||
max_tokens: Maximum tokens to generate
|
||||
"""
|
||||
|
||||
provider: str = "ollama"
|
||||
model: str = "gpt-oss"
|
||||
name: str = "gpt-oss"
|
||||
|
||||
enable_thinking: bool | None = None
|
||||
temperature: float | None = None
|
||||
|
|
@ -47,7 +47,7 @@ class EmbeddingsConfig(BaseModel):
|
|||
model: ModelConfig = Field(
|
||||
default_factory=lambda: ModelConfig(
|
||||
provider="ollama",
|
||||
model="qwen3-embedding:4b",
|
||||
name="qwen3-embedding:4b",
|
||||
)
|
||||
)
|
||||
vector_dim: int = 2560
|
||||
|
|
@ -61,7 +61,7 @@ class QAConfig(BaseModel):
|
|||
model: ModelConfig = Field(
|
||||
default_factory=lambda: ModelConfig(
|
||||
provider="ollama",
|
||||
model="gpt-oss",
|
||||
name="gpt-oss",
|
||||
enable_thinking=False,
|
||||
)
|
||||
)
|
||||
|
|
@ -74,7 +74,7 @@ class ResearchConfig(BaseModel):
|
|||
model: ModelConfig = Field(
|
||||
default_factory=lambda: ModelConfig(
|
||||
provider="ollama",
|
||||
model="gpt-oss",
|
||||
name="gpt-oss",
|
||||
enable_thinking=True,
|
||||
)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -16,7 +16,7 @@ def get_embedder(config: AppConfig = Config) -> EmbedderBase:
|
|||
|
||||
if config.embeddings.model.provider == "ollama":
|
||||
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":
|
||||
|
|
@ -29,21 +29,21 @@ def get_embedder(config: AppConfig = Config) -> EmbedderBase:
|
|||
"uv pip install haiku.rag[voyageai]"
|
||||
)
|
||||
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":
|
||||
from haiku.rag.embeddings.openai import Embedder as 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":
|
||||
from haiku.rag.embeddings.vllm import Embedder as VllmEmbedder
|
||||
|
||||
return VllmEmbedder(
|
||||
config.embeddings.model.model, config.embeddings.vector_dim, config
|
||||
config.embeddings.model.name, config.embeddings.vector_dim, config
|
||||
)
|
||||
|
||||
raise ValueError(
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ from haiku.rag.config import AppConfig, Config
|
|||
|
||||
|
||||
class EmbedderBase:
|
||||
_model: str = Config.embeddings.model.model
|
||||
_model: str = Config.embeddings.model.name
|
||||
_vector_dim: int = Config.embeddings.vector_dim
|
||||
_config: AppConfig = Config
|
||||
|
||||
|
|
|
|||
|
|
@ -45,7 +45,7 @@ def get_reranker(config: AppConfig = Config) -> RerankerBase | None:
|
|||
try:
|
||||
from haiku.rag.reranking.vllm import VLLMReranker
|
||||
|
||||
reranker = VLLMReranker(config.reranking.model.model)
|
||||
reranker = VLLMReranker(config.reranking.model.name)
|
||||
except ImportError:
|
||||
reranker = None
|
||||
|
||||
|
|
@ -54,7 +54,7 @@ def get_reranker(config: AppConfig = Config) -> RerankerBase | None:
|
|||
from haiku.rag.reranking.zeroentropy import ZeroEntropyReranker
|
||||
|
||||
# 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)
|
||||
except ImportError:
|
||||
reranker = None
|
||||
|
|
|
|||
|
|
@ -3,9 +3,7 @@ from haiku.rag.store.models.chunk import Chunk
|
|||
|
||||
|
||||
class RerankerBase:
|
||||
_model: str | None = (
|
||||
Config.reranking.model.model if Config.reranking.model else None
|
||||
)
|
||||
_model: str | None = Config.reranking.model.name if Config.reranking.model else None
|
||||
|
||||
async def rerank(
|
||||
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):
|
||||
def __init__(self):
|
||||
model_name = (
|
||||
Config.reranking.model.model
|
||||
Config.reranking.model.name
|
||||
if Config.reranking.model
|
||||
else "mxbai-rerank-base-v2"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -65,7 +65,7 @@ def get_model(
|
|||
app_config = Config
|
||||
|
||||
provider = model_config.provider
|
||||
model = model_config.model
|
||||
model = model_config.name
|
||||
|
||||
if provider == "ollama":
|
||||
model_settings = None
|
||||
|
|
@ -366,13 +366,13 @@ def prefetch_models():
|
|||
# Collect Ollama models from config
|
||||
required_models: set[str] = set()
|
||||
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":
|
||||
required_models.add(Config.qa.model.model)
|
||||
required_models.add(Config.qa.model.name)
|
||||
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":
|
||||
required_models.add(Config.reranking.model.model)
|
||||
required_models.add(Config.reranking.model.name)
|
||||
|
||||
if not required_models:
|
||||
return
|
||||
|
|
|
|||
|
|
@ -17,7 +17,7 @@ def test_embedder_uses_config_from_get_embedder():
|
|||
embeddings=EmbeddingsConfig(
|
||||
model=ModelConfig(
|
||||
provider="ollama",
|
||||
model="custom-model",
|
||||
name="custom-model",
|
||||
),
|
||||
vector_dim=512,
|
||||
),
|
||||
|
|
@ -42,7 +42,7 @@ def test_vllm_embedder_uses_config():
|
|||
embeddings=EmbeddingsConfig(
|
||||
model=ModelConfig(
|
||||
provider="vllm",
|
||||
model="custom-vllm-model",
|
||||
name="custom-vllm-model",
|
||||
),
|
||||
vector_dim=768,
|
||||
),
|
||||
|
|
@ -68,7 +68,7 @@ def test_openai_embedder_uses_config():
|
|||
embeddings=EmbeddingsConfig(
|
||||
model=ModelConfig(
|
||||
provider="openai",
|
||||
model="text-embedding-3-large",
|
||||
name="text-embedding-3-large",
|
||||
),
|
||||
vector_dim=3072,
|
||||
),
|
||||
|
|
@ -90,7 +90,7 @@ def test_voyageai_embedder_uses_config():
|
|||
embeddings=EmbeddingsConfig(
|
||||
model=ModelConfig(
|
||||
provider="voyageai",
|
||||
model="voyage-large-2",
|
||||
name="voyage-large-2",
|
||||
),
|
||||
vector_dim=1536,
|
||||
),
|
||||
|
|
|
|||
|
|
@ -19,7 +19,7 @@ async def test_qa_ollama(qa_corpus: Dataset, temp_db_path):
|
|||
"""Test Ollama QA with LLM judge."""
|
||||
client = HaikuRAG(temp_db_path)
|
||||
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()
|
||||
|
||||
|
|
@ -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):
|
||||
"""Test OpenAI QA with LLM judge."""
|
||||
client = HaikuRAG(temp_db_path)
|
||||
qa = QuestionAnswerAgent(
|
||||
client, ModelConfig(provider="openai", model="gpt-4o-mini")
|
||||
)
|
||||
qa = QuestionAnswerAgent(client, ModelConfig(provider="openai", name="gpt-4o-mini"))
|
||||
llm_judge = LLMJudge()
|
||||
|
||||
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."""
|
||||
client = HaikuRAG(temp_db_path)
|
||||
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()
|
||||
|
||||
|
|
@ -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):
|
||||
"""Test vLLM QA with LLM judge."""
|
||||
client = HaikuRAG(temp_db_path)
|
||||
qa = QuestionAnswerAgent(
|
||||
client, ModelConfig(provider="vllm", model="Qwen/Qwen3-4B")
|
||||
)
|
||||
qa = QuestionAnswerAgent(client, ModelConfig(provider="vllm", name="Qwen/Qwen3-4B"))
|
||||
llm_judge = LLMJudge()
|
||||
|
||||
doc = qa_corpus[1]
|
||||
|
|
|
|||
|
|
@ -44,7 +44,7 @@ async def test_mxbai_reranker():
|
|||
from haiku.rag.reranking.mxbai import MxBAIReranker
|
||||
|
||||
Config.reranking.model = ModelConfig(
|
||||
provider="mxbai", model="mixedbread-ai/mxbai-rerank-base-v2"
|
||||
provider="mxbai", name="mixedbread-ai/mxbai-rerank-base-v2"
|
||||
)
|
||||
reranker = MxBAIReranker()
|
||||
reranked = await reranker.rerank(
|
||||
|
|
|
|||
|
|
@ -136,16 +136,14 @@ Emoji test: 🚀 ✅ 📝"""
|
|||
|
||||
def test_get_model_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)
|
||||
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."""
|
||||
model_config = ModelConfig(
|
||||
provider="ollama", model="gpt-oss", enable_thinking=False
|
||||
)
|
||||
model_config = ModelConfig(provider="ollama", name="gpt-oss", enable_thinking=False)
|
||||
result = get_model(model_config)
|
||||
assert isinstance(result, OpenAIChatModel)
|
||||
|
||||
|
|
@ -153,7 +151,7 @@ def test_get_model_ollama_with_thinking():
|
|||
def test_get_model_ollama_with_settings():
|
||||
"""Test get_model applies temperature and max_tokens for Ollama."""
|
||||
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)
|
||||
assert isinstance(result, OpenAIChatModel)
|
||||
|
|
@ -161,14 +159,14 @@ def test_get_model_ollama_with_settings():
|
|||
|
||||
def test_get_model_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)
|
||||
assert isinstance(result, OpenAIChatModel)
|
||||
|
||||
|
||||
def test_get_model_openai_with_thinking():
|
||||
"""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)
|
||||
assert isinstance(result, OpenAIChatModel)
|
||||
|
||||
|
|
@ -178,7 +176,7 @@ def test_get_model_anthropic():
|
|||
"""Test get_model returns AnthropicModel for Anthropic."""
|
||||
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)
|
||||
assert isinstance(result, AnthropicModel)
|
||||
|
||||
|
|
@ -190,7 +188,7 @@ def test_get_model_anthropic_with_thinking():
|
|||
|
||||
model_config = ModelConfig(
|
||||
provider="anthropic",
|
||||
model="claude-3-5-sonnet-20241022",
|
||||
name="claude-3-5-sonnet-20241022",
|
||||
enable_thinking=True,
|
||||
)
|
||||
result = get_model(model_config)
|
||||
|
|
@ -202,7 +200,7 @@ def test_get_model_gemini():
|
|||
"""Test get_model returns GoogleModel for Gemini."""
|
||||
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)
|
||||
assert isinstance(result, GoogleModel)
|
||||
|
||||
|
|
@ -213,7 +211,7 @@ def test_get_model_gemini_with_thinking():
|
|||
from pydantic_ai.models.google import GoogleModel
|
||||
|
||||
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)
|
||||
assert isinstance(result, GoogleModel)
|
||||
|
|
@ -224,7 +222,7 @@ def test_get_model_groq():
|
|||
"""Test get_model returns GroqModel for Groq."""
|
||||
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)
|
||||
assert isinstance(result, GroqModel)
|
||||
|
||||
|
|
@ -235,7 +233,7 @@ def test_get_model_groq_with_thinking():
|
|||
from pydantic_ai.models.groq import GroqModel
|
||||
|
||||
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)
|
||||
assert isinstance(result, GroqModel)
|
||||
|
|
@ -247,7 +245,7 @@ def test_get_model_bedrock():
|
|||
from pydantic_ai.models.bedrock import BedrockConverseModel
|
||||
|
||||
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)
|
||||
assert isinstance(result, BedrockConverseModel)
|
||||
|
|
@ -260,7 +258,7 @@ def test_get_model_bedrock_with_thinking():
|
|||
|
||||
model_config = ModelConfig(
|
||||
provider="bedrock",
|
||||
model="anthropic.claude-3-5-sonnet-20241022-v2:0",
|
||||
name="anthropic.claude-3-5-sonnet-20241022-v2:0",
|
||||
enable_thinking=True,
|
||||
)
|
||||
result = get_model(model_config)
|
||||
|
|
@ -269,21 +267,21 @@ def test_get_model_bedrock_with_thinking():
|
|||
|
||||
def test_get_model_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)
|
||||
assert isinstance(result, OpenAIChatModel)
|
||||
|
||||
|
||||
def test_get_model_vllm_with_thinking():
|
||||
"""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)
|
||||
assert isinstance(result, OpenAIChatModel)
|
||||
|
||||
|
||||
def test_get_model_unknown_provider():
|
||||
"""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)
|
||||
assert isinstance(result, str)
|
||||
assert result == "mistral:mistral-large-latest"
|
||||
|
|
@ -293,7 +291,7 @@ def test_get_model_with_all_settings():
|
|||
"""Test get_model applies all settings together."""
|
||||
model_config = ModelConfig(
|
||||
provider="openai",
|
||||
model="gpt-4o",
|
||||
name="gpt-4o",
|
||||
enable_thinking=False,
|
||||
temperature=0.7,
|
||||
max_tokens=500,
|
||||
|
|
|
|||
Loading…
Reference in a new issue