c889a57b6b
Test Suites / Build CI Environment (push) Has been cancelled
Test Suites / Basic Tests (push) Has been cancelled
Test Suites / End-to-End Tests (push) Has been cancelled
Test Suites / CLI Tests (push) Has been cancelled
Test Suites / Slow End-to-End Tests (push) Has been cancelled
Test Suites / Graph Database Tests (push) Has been cancelled
Test Suites / Vector DB Tests (push) Has been cancelled
Test Suites / Temporal Graph Test (push) Has been cancelled
Test Suites / Search Test on Different DBs (push) Has been cancelled
Test Suites / Example Tests (push) Has been cancelled
Test Suites / Notebook Tests (push) Has been cancelled
Test Suites / OS and Python Tests Ubuntu (push) Has been cancelled
Test Suites / OS and Python Tests Extended (push) Has been cancelled
Test Suites / LLM Test Suite (push) Has been cancelled
Test Suites / S3 File Storage Test (push) Has been cancelled
Test Suites / Run Integration Tests (push) Has been cancelled
Test Suites / MCP Tests (push) Has been cancelled
Test Suites / Docker Compose Test (push) Has been cancelled
Test Suites / Docker CI test (push) Has been cancelled
Test Suites / Relational DB Migration Tests (push) Has been cancelled
Test Suites / Distributed Cognee Test (push) Has been cancelled
Test Suites / DB Examples Tests (push) Has been cancelled
Test Suites / Test Completion Status (push) Has been cancelled
Test Suites / Claude Code Review (push) Has been cancelled
Test Suites / basic checks (push) Has been cancelled
build | Build and Push Cognee MCP Docker Image to dockerhub / docker-build-and-push (push) Has been cancelled
Scorecard supply-chain security / Scorecard analysis (push) Has been cancelled
build | Build and Push Docker Image to dockerhub / docker-build-and-push (push) Has been cancelled
Weighted Edges Tests / Test Weighted Edges Core Functionality (3.11) (push) Has been cancelled
Weighted Edges Tests / Test Weighted Edges Core Functionality (3.12) (push) Has been cancelled
Weighted Edges Tests / Test Weighted Edges with Different Graph Databases (kuzu, kuzu) (push) Has been cancelled
Weighted Edges Tests / Test Weighted Edges with Different Graph Databases (neo4j, neo4j) (push) Has been cancelled
Weighted Edges Tests / Test Weighted Edges Examples (push) Has been cancelled
Weighted Edges Tests / Code Quality for Weighted Edges (push) Has been cancelled
108 lines
4.0 KiB
Python
108 lines
4.0 KiB
Python
"""
|
|
Tests for OpenAICompatibleEmbeddingEngine.
|
|
|
|
Verifies that the engine:
|
|
- Returns mock embeddings when MOCK_EMBEDDING is set
|
|
- Calls the OpenAI SDK with encoding_format="float"
|
|
- Reports correct vector size and batch size
|
|
"""
|
|
|
|
from unittest.mock import AsyncMock, MagicMock, patch
|
|
|
|
import pytest
|
|
|
|
|
|
class TestOpenAICompatibleEmbeddingEngine:
|
|
"""Unit tests for OpenAICompatibleEmbeddingEngine."""
|
|
|
|
def _make_engine(self, **kwargs):
|
|
"""Create an engine instance with defaults suitable for testing."""
|
|
defaults = {
|
|
"model": "test-model",
|
|
"dimensions": 4096,
|
|
"max_completion_tokens": 8191,
|
|
"endpoint": "http://localhost:8099",
|
|
"api_key": "test-key",
|
|
"batch_size": 36,
|
|
}
|
|
defaults.update(kwargs)
|
|
|
|
from cognee.infrastructure.databases.vector.embeddings.OpenAICompatibleEmbeddingEngine import (
|
|
OpenAICompatibleEmbeddingEngine,
|
|
)
|
|
|
|
return OpenAICompatibleEmbeddingEngine(**defaults)
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_mock_embedding(self, monkeypatch):
|
|
"""When MOCK_EMBEDDING=true, embed_text returns zero vectors of correct dimensions."""
|
|
monkeypatch.setenv("MOCK_EMBEDDING", "true")
|
|
engine = self._make_engine(dimensions=4096)
|
|
result = await engine.embed_text(["hello", "world"])
|
|
assert len(result) == 2
|
|
assert len(result[0]) == 4096
|
|
assert all(v == 0.0 for v in result[0])
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_embed_text_calls_openai_with_encoding_format_float(self, monkeypatch):
|
|
"""embed_text must call OpenAI SDK with encoding_format='float'."""
|
|
monkeypatch.delenv("MOCK_EMBEDDING", raising=False)
|
|
|
|
engine = self._make_engine()
|
|
|
|
# Build a mock response matching OpenAI SDK's CreateEmbeddingResponse
|
|
mock_item = MagicMock()
|
|
mock_item.embedding = [0.1] * 4096
|
|
|
|
mock_response = MagicMock()
|
|
mock_response.data = [mock_item]
|
|
|
|
# Mock the AsyncOpenAI client's embeddings.create
|
|
engine._client = MagicMock()
|
|
engine._client.embeddings.create = AsyncMock(return_value=mock_response)
|
|
|
|
result = await engine.embed_text(["test text"])
|
|
|
|
# Verify create was called with encoding_format="float"
|
|
engine._client.embeddings.create.assert_called_once_with(
|
|
model="test-model",
|
|
input=["test text"],
|
|
encoding_format="float",
|
|
)
|
|
|
|
assert len(result) == 1
|
|
assert len(result[0]) == 4096
|
|
|
|
def test_get_vector_size(self):
|
|
"""get_vector_size returns the configured dimensions."""
|
|
engine = self._make_engine(dimensions=768)
|
|
assert engine.get_vector_size() == 768
|
|
|
|
def test_get_batch_size(self):
|
|
"""get_batch_size returns the configured batch size."""
|
|
engine = self._make_engine(batch_size=50)
|
|
assert engine.get_batch_size() == 50
|
|
|
|
def test_max_completion_tokens_is_exposed(self):
|
|
"""The engine exposes max_completion_tokens for chunk sizing logic."""
|
|
engine = self._make_engine(max_completion_tokens=2048)
|
|
assert engine.max_completion_tokens == 2048
|
|
|
|
def test_endpoint_normalization(self):
|
|
"""Endpoint without /v1 gets /v1 appended for the SDK base_url."""
|
|
engine = self._make_engine(endpoint="http://localhost:8099")
|
|
assert str(engine._client._base_url).rstrip("/").endswith("/v1")
|
|
|
|
engine2 = self._make_engine(endpoint="http://localhost:8099/v1")
|
|
assert str(engine2._client._base_url).rstrip("/").endswith("/v1")
|
|
|
|
# Both should produce equivalent normalized URLs
|
|
assert str(engine._client._base_url) == str(engine2._client._base_url)
|
|
|
|
def test_endpoint_normalization_strips_embeddings_suffix(self):
|
|
"""Endpoint with /v1/embeddings should not produce /v1/embeddings/v1."""
|
|
engine = self._make_engine(endpoint="http://localhost:8099/v1/embeddings")
|
|
base_url = str(engine._client._base_url).rstrip("/")
|
|
assert base_url.endswith("/v1")
|
|
assert "/embeddings" not in base_url
|