Update tests and client for explicit embed_query/embed_documents API
This commit is contained in:
parent
132b8a36bc
commit
3852e961b9
4 changed files with 93 additions and 140 deletions
|
|
@ -1407,7 +1407,7 @@ class HaikuRAG:
|
||||||
|
|
||||||
# Generate new embeddings using contextualize for consistency
|
# Generate new embeddings using contextualize for consistency
|
||||||
texts = contextualize(chunks)
|
texts = contextualize(chunks)
|
||||||
embeddings = await self.chunk_repository.embedder.embed(texts)
|
embeddings = await self.chunk_repository.embedder.embed_documents(texts)
|
||||||
|
|
||||||
# Build updated records
|
# Build updated records
|
||||||
for chunk, embedding in zip(chunks, embeddings):
|
for chunk, embedding in zip(chunks, embeddings):
|
||||||
|
|
|
||||||
|
|
@ -3,20 +3,16 @@ import os
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from haiku.rag.config import Config
|
from haiku.rag.config import AppConfig, EmbeddingModelConfig, EmbeddingsConfig
|
||||||
from haiku.rag.embeddings import contextualize, embed_chunks
|
from haiku.rag.embeddings import contextualize, embed_chunks, get_embedder
|
||||||
from haiku.rag.embeddings.ollama import Embedder as OllamaEmbedder
|
|
||||||
from haiku.rag.embeddings.openai import Embedder as OpenAIEmbedder
|
|
||||||
from haiku.rag.embeddings.vllm import Embedder as VLLMEmbedder
|
|
||||||
from haiku.rag.store.models.chunk import Chunk
|
from haiku.rag.store.models.chunk import Chunk
|
||||||
|
|
||||||
OPENAI_AVAILABLE = bool(os.getenv("OPENAI_API_KEY"))
|
OPENAI_AVAILABLE = bool(os.getenv("OPENAI_API_KEY"))
|
||||||
VOYAGEAI_AVAILABLE = bool(os.getenv("VOYAGE_API_KEY"))
|
VOYAGEAI_AVAILABLE = bool(os.getenv("VOYAGE_API_KEY"))
|
||||||
VLLM_EMBEDDINGS_AVAILABLE = bool(Config.providers.vllm.embeddings_base_url)
|
|
||||||
|
|
||||||
|
|
||||||
# Calculate cosine similarity
|
|
||||||
def similarities(embeddings, test_embedding):
|
def similarities(embeddings, test_embedding):
|
||||||
|
"""Calculate cosine similarity between embeddings and a test embedding."""
|
||||||
return [
|
return [
|
||||||
np.dot(embedding, test_embedding)
|
np.dot(embedding, test_embedding)
|
||||||
/ (np.linalg.norm(embedding) * np.linalg.norm(test_embedding))
|
/ (np.linalg.norm(embedding) * np.linalg.norm(test_embedding))
|
||||||
|
|
@ -26,35 +22,41 @@ def similarities(embeddings, test_embedding):
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_ollama_embedder():
|
async def test_ollama_embedder():
|
||||||
embedder = OllamaEmbedder("mxbai-embed-large", 1024)
|
"""Test Ollama embedder via pydantic-ai."""
|
||||||
|
config = AppConfig(
|
||||||
|
embeddings=EmbeddingsConfig(
|
||||||
|
model=EmbeddingModelConfig(
|
||||||
|
provider="ollama", name="mxbai-embed-large", vector_dim=1024
|
||||||
|
)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
embedder = get_embedder(config)
|
||||||
phrases = [
|
phrases = [
|
||||||
"I enjoy eating great food.",
|
"I enjoy eating great food.",
|
||||||
"Python is my favorite programming language.",
|
"Python is my favorite programming language.",
|
||||||
"I love to travel and see new places.",
|
"I love to travel and see new places.",
|
||||||
]
|
]
|
||||||
|
|
||||||
# Test batch embedding
|
# Test batch embedding (documents)
|
||||||
embeddings = await embedder.embed(phrases)
|
embeddings = await embedder.embed_documents(phrases)
|
||||||
assert isinstance(embeddings, list)
|
assert isinstance(embeddings, list)
|
||||||
assert len(embeddings) == 3
|
assert len(embeddings) == 3
|
||||||
assert all(isinstance(emb, list) for emb in embeddings)
|
assert all(isinstance(emb, list) for emb in embeddings)
|
||||||
embeddings = [np.array(emb) for emb in embeddings]
|
embeddings = [np.array(emb) for emb in embeddings]
|
||||||
|
|
||||||
|
# Test query embedding
|
||||||
test_phrase = "I am going for a camping trip."
|
test_phrase = "I am going for a camping trip."
|
||||||
test_embedding = await embedder.embed(test_phrase)
|
test_embedding = await embedder.embed_query(test_phrase)
|
||||||
|
|
||||||
sims = similarities(embeddings, test_embedding)
|
sims = similarities(embeddings, test_embedding)
|
||||||
assert max(sims) == sims[2]
|
assert max(sims) == sims[2]
|
||||||
|
|
||||||
test_phrase = "When is dinner ready?"
|
test_phrase = "When is dinner ready?"
|
||||||
test_embedding = await embedder.embed(test_phrase)
|
test_embedding = await embedder.embed_query(test_phrase)
|
||||||
|
|
||||||
sims = similarities(embeddings, test_embedding)
|
sims = similarities(embeddings, test_embedding)
|
||||||
assert max(sims) == sims[0]
|
assert max(sims) == sims[0]
|
||||||
|
|
||||||
test_phrase = "I work as a software developer."
|
test_phrase = "I work as a software developer."
|
||||||
test_embedding = await embedder.embed(test_phrase)
|
test_embedding = await embedder.embed_query(test_phrase)
|
||||||
|
|
||||||
sims = similarities(embeddings, test_embedding)
|
sims = similarities(embeddings, test_embedding)
|
||||||
assert max(sims) == sims[1]
|
assert max(sims) == sims[1]
|
||||||
|
|
||||||
|
|
@ -62,35 +64,41 @@ async def test_ollama_embedder():
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@pytest.mark.skipif(not OPENAI_AVAILABLE, reason="OpenAI API key not available")
|
@pytest.mark.skipif(not OPENAI_AVAILABLE, reason="OpenAI API key not available")
|
||||||
async def test_openai_embedder():
|
async def test_openai_embedder():
|
||||||
embedder = OpenAIEmbedder("text-embedding-3-small", 1536)
|
"""Test OpenAI embedder via pydantic-ai."""
|
||||||
|
config = AppConfig(
|
||||||
|
embeddings=EmbeddingsConfig(
|
||||||
|
model=EmbeddingModelConfig(
|
||||||
|
provider="openai", name="text-embedding-3-small", vector_dim=1536
|
||||||
|
)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
embedder = get_embedder(config)
|
||||||
phrases = [
|
phrases = [
|
||||||
"I enjoy eating great food.",
|
"I enjoy eating great food.",
|
||||||
"Python is my favorite programming language.",
|
"Python is my favorite programming language.",
|
||||||
"I love to travel and see new places.",
|
"I love to travel and see new places.",
|
||||||
]
|
]
|
||||||
|
|
||||||
# Test batch embedding
|
# Test batch embedding (documents)
|
||||||
embeddings = await embedder.embed(phrases)
|
embeddings = await embedder.embed_documents(phrases)
|
||||||
assert isinstance(embeddings, list)
|
assert isinstance(embeddings, list)
|
||||||
assert len(embeddings) == 3
|
assert len(embeddings) == 3
|
||||||
assert all(isinstance(emb, list) for emb in embeddings)
|
assert all(isinstance(emb, list) for emb in embeddings)
|
||||||
embeddings = [np.array(emb) for emb in embeddings]
|
embeddings = [np.array(emb) for emb in embeddings]
|
||||||
|
|
||||||
|
# Test query embedding
|
||||||
test_phrase = "I am going for a camping trip."
|
test_phrase = "I am going for a camping trip."
|
||||||
test_embedding = await embedder.embed(test_phrase)
|
test_embedding = await embedder.embed_query(test_phrase)
|
||||||
|
|
||||||
sims = similarities(embeddings, test_embedding)
|
sims = similarities(embeddings, test_embedding)
|
||||||
assert max(sims) == sims[2]
|
assert max(sims) == sims[2]
|
||||||
|
|
||||||
test_phrase = "When is dinner ready?"
|
test_phrase = "When is dinner ready?"
|
||||||
test_embedding = await embedder.embed(test_phrase)
|
test_embedding = await embedder.embed_query(test_phrase)
|
||||||
|
|
||||||
sims = similarities(embeddings, test_embedding)
|
sims = similarities(embeddings, test_embedding)
|
||||||
assert max(sims) == sims[0]
|
assert max(sims) == sims[0]
|
||||||
|
|
||||||
test_phrase = "I work as a software developer."
|
test_phrase = "I work as a software developer."
|
||||||
test_embedding = await embedder.embed(test_phrase)
|
test_embedding = await embedder.embed_query(test_phrase)
|
||||||
|
|
||||||
sims = similarities(embeddings, test_embedding)
|
sims = similarities(embeddings, test_embedding)
|
||||||
assert max(sims) == sims[1]
|
assert max(sims) == sims[1]
|
||||||
|
|
||||||
|
|
@ -98,38 +106,42 @@ async def test_openai_embedder():
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@pytest.mark.skipif(not VOYAGEAI_AVAILABLE, reason="VoyageAI API key not available")
|
@pytest.mark.skipif(not VOYAGEAI_AVAILABLE, reason="VoyageAI API key not available")
|
||||||
async def test_voyageai_embedder():
|
async def test_voyageai_embedder():
|
||||||
|
"""Test VoyageAI embedder."""
|
||||||
try:
|
try:
|
||||||
from haiku.rag.embeddings.voyageai import Embedder as VoyageAIEmbedder
|
config = AppConfig(
|
||||||
|
embeddings=EmbeddingsConfig(
|
||||||
embedder = VoyageAIEmbedder("voyage-3.5", 1024)
|
model=EmbeddingModelConfig(
|
||||||
|
provider="voyageai", name="voyage-3.5", vector_dim=1024
|
||||||
|
)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
embedder = get_embedder(config)
|
||||||
phrases = [
|
phrases = [
|
||||||
"I enjoy eating great food.",
|
"I enjoy eating great food.",
|
||||||
"Python is my favorite programming language.",
|
"Python is my favorite programming language.",
|
||||||
"I love to travel and see new places.",
|
"I love to travel and see new places.",
|
||||||
]
|
]
|
||||||
|
|
||||||
# Test batch embedding
|
# Test batch embedding (documents)
|
||||||
embeddings = await embedder.embed(phrases)
|
embeddings = await embedder.embed_documents(phrases)
|
||||||
assert isinstance(embeddings, list)
|
assert isinstance(embeddings, list)
|
||||||
assert len(embeddings) == 3
|
assert len(embeddings) == 3
|
||||||
assert all(isinstance(emb, list) for emb in embeddings)
|
assert all(isinstance(emb, list) for emb in embeddings)
|
||||||
embeddings = [np.array(emb) for emb in embeddings]
|
embeddings = [np.array(emb) for emb in embeddings]
|
||||||
|
|
||||||
|
# Test query embedding
|
||||||
test_phrase = "I am going for a camping trip."
|
test_phrase = "I am going for a camping trip."
|
||||||
test_embedding = await embedder.embed(test_phrase)
|
test_embedding = await embedder.embed_query(test_phrase)
|
||||||
|
|
||||||
sims = similarities(embeddings, test_embedding)
|
sims = similarities(embeddings, test_embedding)
|
||||||
assert max(sims) == sims[2]
|
assert max(sims) == sims[2]
|
||||||
|
|
||||||
test_phrase = "When is dinner ready?"
|
test_phrase = "When is dinner ready?"
|
||||||
test_embedding = await embedder.embed(test_phrase)
|
test_embedding = await embedder.embed_query(test_phrase)
|
||||||
|
|
||||||
sims = similarities(embeddings, test_embedding)
|
sims = similarities(embeddings, test_embedding)
|
||||||
assert max(sims) == sims[0]
|
assert max(sims) == sims[0]
|
||||||
|
|
||||||
test_phrase = "I work as a software developer."
|
test_phrase = "I work as a software developer."
|
||||||
test_embedding = await embedder.embed(test_phrase)
|
test_embedding = await embedder.embed_query(test_phrase)
|
||||||
|
|
||||||
sims = similarities(embeddings, test_embedding)
|
sims = similarities(embeddings, test_embedding)
|
||||||
assert max(sims) == sims[1]
|
assert max(sims) == sims[1]
|
||||||
|
|
||||||
|
|
@ -137,44 +149,6 @@ async def test_voyageai_embedder():
|
||||||
pytest.skip("VoyageAI package not installed")
|
pytest.skip("VoyageAI package not installed")
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
@pytest.mark.skipif(
|
|
||||||
not VLLM_EMBEDDINGS_AVAILABLE, reason="vLLM embeddings server not configured"
|
|
||||||
)
|
|
||||||
async def test_vllm_embedder():
|
|
||||||
embedder = VLLMEmbedder("mixedbread-ai/mxbai-embed-large-v1", 512)
|
|
||||||
phrases = [
|
|
||||||
"I enjoy eating great food.",
|
|
||||||
"Python is my favorite programming language.",
|
|
||||||
"I love to travel and see new places.",
|
|
||||||
]
|
|
||||||
|
|
||||||
# Test batch embedding
|
|
||||||
embeddings = await embedder.embed(phrases)
|
|
||||||
assert isinstance(embeddings, list)
|
|
||||||
assert len(embeddings) == 3
|
|
||||||
assert all(isinstance(emb, list) for emb in embeddings)
|
|
||||||
embeddings = [np.array(emb) for emb in embeddings]
|
|
||||||
|
|
||||||
test_phrase = "I am going for a camping trip."
|
|
||||||
test_embedding = await embedder.embed(test_phrase)
|
|
||||||
|
|
||||||
sims = similarities(embeddings, test_embedding)
|
|
||||||
assert max(sims) == sims[2]
|
|
||||||
|
|
||||||
test_phrase = "When is dinner ready?"
|
|
||||||
test_embedding = await embedder.embed(test_phrase)
|
|
||||||
|
|
||||||
sims = similarities(embeddings, test_embedding)
|
|
||||||
assert max(sims) == sims[0]
|
|
||||||
|
|
||||||
test_phrase = "I work as a software developer."
|
|
||||||
test_embedding = await embedder.embed(test_phrase)
|
|
||||||
|
|
||||||
sims = similarities(embeddings, test_embedding)
|
|
||||||
assert max(sims) == sims[1]
|
|
||||||
|
|
||||||
|
|
||||||
def test_contextualize_with_headings():
|
def test_contextualize_with_headings():
|
||||||
"""Test that contextualize prepends headings to chunk content."""
|
"""Test that contextualize prepends headings to chunk content."""
|
||||||
chunks = [
|
chunks = [
|
||||||
|
|
|
||||||
|
|
@ -4,16 +4,14 @@ from haiku.rag.config import (
|
||||||
AppConfig,
|
AppConfig,
|
||||||
EmbeddingModelConfig,
|
EmbeddingModelConfig,
|
||||||
EmbeddingsConfig,
|
EmbeddingsConfig,
|
||||||
LMStudioConfig,
|
|
||||||
OllamaConfig,
|
OllamaConfig,
|
||||||
ProvidersConfig,
|
ProvidersConfig,
|
||||||
VLLMConfig,
|
|
||||||
)
|
)
|
||||||
from haiku.rag.embeddings import get_embedder
|
from haiku.rag.embeddings import get_embedder
|
||||||
|
|
||||||
|
|
||||||
def test_embedder_uses_config_from_get_embedder():
|
def test_ollama_embedder_uses_config():
|
||||||
"""Test that embedders use the config passed to get_embedder."""
|
"""Test that Ollama embedder uses the config passed to get_embedder."""
|
||||||
custom_config = AppConfig(
|
custom_config = AppConfig(
|
||||||
embeddings=EmbeddingsConfig(
|
embeddings=EmbeddingsConfig(
|
||||||
model=EmbeddingModelConfig(
|
model=EmbeddingModelConfig(
|
||||||
|
|
@ -22,41 +20,16 @@ def test_embedder_uses_config_from_get_embedder():
|
||||||
),
|
),
|
||||||
providers=ProvidersConfig(
|
providers=ProvidersConfig(
|
||||||
ollama=OllamaConfig(base_url="http://custom-ollama:8080"),
|
ollama=OllamaConfig(base_url="http://custom-ollama:8080"),
|
||||||
vllm=VLLMConfig(embeddings_base_url="http://custom-vllm:9000"),
|
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
embedder = get_embedder(custom_config)
|
embedder = get_embedder(custom_config)
|
||||||
|
|
||||||
assert embedder._model == "custom-model"
|
|
||||||
assert embedder._vector_dim == 512
|
assert embedder._vector_dim == 512
|
||||||
assert embedder._config.providers.ollama.base_url == "http://custom-ollama:8080"
|
|
||||||
|
|
||||||
|
|
||||||
def test_vllm_embedder_uses_config():
|
|
||||||
"""Test that vllm embedder uses the config passed to get_embedder."""
|
|
||||||
custom_config = AppConfig(
|
|
||||||
embeddings=EmbeddingsConfig(
|
|
||||||
model=EmbeddingModelConfig(
|
|
||||||
provider="vllm", name="custom-vllm-model", vector_dim=768
|
|
||||||
),
|
|
||||||
),
|
|
||||||
providers=ProvidersConfig(
|
|
||||||
vllm=VLLMConfig(embeddings_base_url="http://custom-vllm:9001"),
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
embedder = get_embedder(custom_config)
|
|
||||||
|
|
||||||
assert embedder._model == "custom-vllm-model"
|
|
||||||
assert embedder._vector_dim == 768
|
|
||||||
assert (
|
|
||||||
embedder._config.providers.vllm.embeddings_base_url == "http://custom-vllm:9001"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
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(
|
||||||
model=EmbeddingModelConfig(
|
model=EmbeddingModelConfig(
|
||||||
|
|
@ -67,48 +40,68 @@ def test_openai_embedder_uses_config():
|
||||||
|
|
||||||
embedder = get_embedder(custom_config)
|
embedder = get_embedder(custom_config)
|
||||||
|
|
||||||
assert embedder._model == "text-embedding-3-large"
|
|
||||||
assert embedder._vector_dim == 3072
|
assert embedder._vector_dim == 3072
|
||||||
assert embedder._config == custom_config
|
|
||||||
|
|
||||||
|
|
||||||
def test_lm_studio_embedder_uses_config():
|
def test_openai_embedder_with_base_url():
|
||||||
"""Test that lm_studio embedder uses the config passed to get_embedder."""
|
"""Test that OpenAI embedder uses custom base_url for vLLM/LM Studio."""
|
||||||
custom_config = AppConfig(
|
custom_config = AppConfig(
|
||||||
embeddings=EmbeddingsConfig(
|
embeddings=EmbeddingsConfig(
|
||||||
model=EmbeddingModelConfig(
|
model=EmbeddingModelConfig(
|
||||||
provider="lm_studio", name="custom-lm-studio-model", vector_dim=1024
|
provider="openai",
|
||||||
|
name="some-local-model",
|
||||||
|
vector_dim=768,
|
||||||
|
base_url="http://localhost:8000/v1",
|
||||||
|
),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
embedder = get_embedder(custom_config)
|
||||||
|
|
||||||
|
assert embedder._vector_dim == 768
|
||||||
|
|
||||||
|
|
||||||
|
def test_cohere_embedder_uses_config():
|
||||||
|
"""Test that Cohere embedder uses the config passed to get_embedder."""
|
||||||
|
custom_config = AppConfig(
|
||||||
|
embeddings=EmbeddingsConfig(
|
||||||
|
model=EmbeddingModelConfig(
|
||||||
|
provider="cohere", name="embed-v4.0", vector_dim=1024
|
||||||
),
|
),
|
||||||
),
|
),
|
||||||
providers=ProvidersConfig(
|
|
||||||
lm_studio=LMStudioConfig(base_url="http://custom-lmstudio:5678"),
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
embedder = get_embedder(custom_config)
|
embedder = get_embedder(custom_config)
|
||||||
|
|
||||||
assert embedder._model == "custom-lm-studio-model"
|
|
||||||
assert embedder._vector_dim == 1024
|
assert embedder._vector_dim == 1024
|
||||||
assert (
|
|
||||||
embedder._config.providers.lm_studio.base_url == "http://custom-lmstudio:5678"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.skipif(
|
def test_sentence_transformers_embedder_uses_config():
|
||||||
True, reason="VoyageAI is an optional dependency, may not be installed"
|
"""Test that SentenceTransformers embedder uses the config."""
|
||||||
)
|
|
||||||
def test_voyageai_embedder_uses_config():
|
|
||||||
"""Test that voyageai embedder uses the config passed to get_embedder."""
|
|
||||||
custom_config = AppConfig(
|
custom_config = AppConfig(
|
||||||
embeddings=EmbeddingsConfig(
|
embeddings=EmbeddingsConfig(
|
||||||
model=EmbeddingModelConfig(
|
model=EmbeddingModelConfig(
|
||||||
provider="voyageai", name="voyage-large-2", vector_dim=1536
|
provider="sentence-transformers",
|
||||||
|
name="all-MiniLM-L6-v2",
|
||||||
|
vector_dim=384,
|
||||||
),
|
),
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
embedder = get_embedder(custom_config)
|
embedder = get_embedder(custom_config)
|
||||||
|
|
||||||
assert embedder._model == "voyage-large-2"
|
assert embedder._vector_dim == 384
|
||||||
assert embedder._vector_dim == 1536
|
|
||||||
assert embedder._config == custom_config
|
|
||||||
|
def test_unsupported_provider_raises():
|
||||||
|
"""Test that unsupported provider raises ValueError."""
|
||||||
|
custom_config = AppConfig(
|
||||||
|
embeddings=EmbeddingsConfig(
|
||||||
|
model=EmbeddingModelConfig(
|
||||||
|
provider="unsupported-provider", name="model", vector_dim=512
|
||||||
|
),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="Unsupported embedding provider"):
|
||||||
|
get_embedder(custom_config)
|
||||||
|
|
|
||||||
|
|
@ -270,20 +270,6 @@ def test_get_model_bedrock_with_thinking():
|
||||||
assert isinstance(result, BedrockConverseModel)
|
assert isinstance(result, BedrockConverseModel)
|
||||||
|
|
||||||
|
|
||||||
def test_get_model_vllm():
|
|
||||||
"""Test get_model returns OpenAIChatModel for vLLM."""
|
|
||||||
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", name="gpt-oss", enable_thinking=False)
|
|
||||||
result = get_model(model_config)
|
|
||||||
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", name="mistral-large-latest")
|
model_config = ModelConfig(provider="mistral", name="mistral-large-latest")
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue