145 lines
4.6 KiB
Python
145 lines
4.6 KiB
Python
"""Tests for schemas module."""
|
|
|
|
import pytest
|
|
from pydantic import ValidationError
|
|
|
|
from rerank.schemas import RerankRequest, RerankResponse
|
|
from rerank.types import BackendType, RerankResult, RerankUsage
|
|
|
|
|
|
class TestRerankRequest:
|
|
"""Tests for RerankRequest schema."""
|
|
|
|
def test_basic_request(self) -> None:
|
|
"""Test basic request creation."""
|
|
request = RerankRequest(
|
|
model="test-model",
|
|
query="What is AI?",
|
|
documents=["doc1", "doc2", "doc3"],
|
|
)
|
|
assert request.model == "test-model"
|
|
assert request.query == "What is AI?"
|
|
assert len(request.documents) == 3
|
|
|
|
def test_empty_documents_rejected(self) -> None:
|
|
"""Test that empty documents list is rejected."""
|
|
with pytest.raises(ValidationError):
|
|
RerankRequest(
|
|
model="test-model",
|
|
query="query",
|
|
documents=[],
|
|
)
|
|
|
|
def test_empty_string_in_documents_rejected(self) -> None:
|
|
"""Test that empty strings in documents are rejected."""
|
|
with pytest.raises(ValidationError):
|
|
RerankRequest(
|
|
model="test-model",
|
|
query="query",
|
|
documents=["doc1", "", "doc3"],
|
|
)
|
|
|
|
def test_empty_query_rejected(self) -> None:
|
|
"""Test that empty query is rejected."""
|
|
with pytest.raises(ValidationError):
|
|
RerankRequest(
|
|
model="test-model",
|
|
query="",
|
|
documents=["doc1", "doc2"],
|
|
)
|
|
|
|
def test_top_n_parameter(self) -> None:
|
|
"""Test top_n parameter."""
|
|
request = RerankRequest(
|
|
model="test-model",
|
|
query="query",
|
|
documents=["doc1", "doc2", "doc3"],
|
|
top_n=2,
|
|
)
|
|
assert request.top_n == 2
|
|
|
|
def test_return_documents_parameter(self) -> None:
|
|
"""Test return_documents parameter."""
|
|
request = RerankRequest(
|
|
model="test-model",
|
|
query="query",
|
|
documents=["doc1", "doc2"],
|
|
return_documents=True,
|
|
)
|
|
assert request.return_documents is True
|
|
|
|
def test_backend_override(self) -> None:
|
|
"""Test backend override."""
|
|
request = RerankRequest(
|
|
model="test-model",
|
|
query="query",
|
|
documents=["doc1", "doc2"],
|
|
backend=BackendType.LLAMACPP,
|
|
)
|
|
assert request.backend == BackendType.LLAMACPP
|
|
|
|
def test_invalid_top_n(self) -> None:
|
|
"""Test that invalid top_n is rejected."""
|
|
with pytest.raises(ValidationError):
|
|
RerankRequest(
|
|
model="test-model",
|
|
query="query",
|
|
documents=["doc1", "doc2"],
|
|
top_n=0,
|
|
)
|
|
|
|
|
|
class TestRerankResponse:
|
|
"""Tests for RerankResponse schema."""
|
|
|
|
def test_creation(self) -> None:
|
|
"""Test creating response."""
|
|
response = RerankResponse(
|
|
model="test-model",
|
|
results=[
|
|
RerankResult(index=2, relevance_score=0.95),
|
|
RerankResult(index=0, relevance_score=0.82),
|
|
],
|
|
usage=RerankUsage(total_tokens=100),
|
|
backend="vllm",
|
|
)
|
|
assert response.model == "test-model"
|
|
assert len(response.results) == 2
|
|
assert response.backend == "vllm"
|
|
assert response.id.startswith("rerank-")
|
|
|
|
def test_results_ordering(self) -> None:
|
|
"""Test that results maintain order."""
|
|
results = [
|
|
RerankResult(index=2, relevance_score=0.95),
|
|
RerankResult(index=0, relevance_score=0.82),
|
|
RerankResult(index=1, relevance_score=0.10),
|
|
]
|
|
response = RerankResponse(
|
|
model="test-model",
|
|
results=results,
|
|
usage=RerankUsage(total_tokens=100),
|
|
backend="vllm",
|
|
)
|
|
assert response.results[0].index == 2
|
|
assert response.results[0].relevance_score == 0.95
|
|
assert response.results[1].index == 0
|
|
assert response.results[2].index == 1
|
|
|
|
|
|
class TestRerankResult:
|
|
"""Tests for RerankResult model."""
|
|
|
|
def test_without_document(self) -> None:
|
|
"""Test result without document."""
|
|
result = RerankResult(index=0, relevance_score=0.95)
|
|
assert result.document is None
|
|
|
|
def test_with_document(self) -> None:
|
|
"""Test result with document."""
|
|
result = RerankResult(
|
|
index=0,
|
|
relevance_score=0.95,
|
|
document="This is the document text",
|
|
)
|
|
assert result.document == "This is the document text"
|