ACCESS-EMBEDDING / tests /test_embed.py
FadQ's picture
feat: embeddinf service foe access"
4398e81
Raw
History Blame Contribute Delete
4.23 kB
import math
import pytest
from fastapi.testclient import TestClient
from app.config import DEFAULT_DIMENSION, get_settings
from app.main import app
class FakeVector:
def __init__(self, values: list[float]) -> None:
self.values = values
def tolist(self) -> list[float]:
return self.values
class FakeMatrix:
def __init__(self, vectors: list[list[float]]) -> None:
self.vectors = vectors
def tolist(self) -> list[list[float]]:
return self.vectors
class FakeModel:
def encode(self, value, normalize_embeddings: bool = True):
if isinstance(value, list):
return FakeMatrix([fake_embedding(item) for item in value])
return FakeVector(fake_embedding(value))
def fake_embedding(text: str) -> list[float]:
lowered = text.lower()
value = 0.01
if any(keyword in lowered for keyword in ["saldo", "pembayaran", "tiket", "e-ticket"]):
value = 0.08
elif any(keyword in lowered for keyword in ["kereta", "terlambat", "stasiun"]):
value = -0.08
vector = [value] * DEFAULT_DIMENSION
norm = math.sqrt(sum(item * item for item in vector))
return [item / norm for item in vector]
@pytest.fixture(autouse=True)
def fake_model(monkeypatch):
from app import embedding
monkeypatch.setattr(embedding, "get_model", lambda: FakeModel())
def cosine_similarity(left: list[float], right: list[float]) -> float:
dot = sum(a * b for a, b in zip(left, right))
left_norm = math.sqrt(sum(value * value for value in left))
right_norm = math.sqrt(sum(value * value for value in right))
return dot / (left_norm * right_norm)
def test_embed_requires_api_key_when_configured(monkeypatch) -> None:
monkeypatch.setenv("EMBEDDING_API_KEY", "test-secret")
get_settings.cache_clear()
client = TestClient(app)
missing_key_response = client.post("/embed", json={"text": "Saldo terpotong."})
wrong_key_response = client.post(
"/embed",
headers={"X-API-Key": "wrong"},
json={"text": "Saldo terpotong."},
)
assert missing_key_response.status_code == 401
assert wrong_key_response.status_code == 401
monkeypatch.delenv("EMBEDDING_API_KEY", raising=False)
get_settings.cache_clear()
def test_embed_rejects_invalid_input() -> None:
client = TestClient(app)
response = client.post("/embed", json={"text": " "})
assert response.status_code == 422
assert response.json()["detail"] == "Text is required"
def test_embed_returns_normalized_vector() -> None:
client = TestClient(app)
response = client.post(
"/embed",
json={"text": "Saldo saya terpotong tapi tiket tidak muncul."},
)
body = response.json()
assert response.status_code == 200
assert body["dimension"] == 384
assert body["model"] == "sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2"
assert len(body["embedding"]) == 384
assert all(isinstance(value, float) for value in body["embedding"])
def test_batch_embed_preserves_input_order() -> None:
client = TestClient(app)
response = client.post(
"/embed/batch",
json={
"texts": [
"Saldo terpotong tapi tiket belum muncul.",
"Refund pembatalan tiket belum diterima.",
"Kereta terlambat tanpa pemberitahuan.",
],
},
)
body = response.json()
assert response.status_code == 200
assert body["dimension"] == 384
assert len(body["embeddings"]) == 3
assert all(len(vector) == 384 for vector in body["embeddings"])
def test_semantic_similarity_sanity() -> None:
client = TestClient(app)
query = client.post(
"/embed",
json={"text": "Saldo saya terpotong tapi tiket tidak muncul."},
).json()["embedding"]
similar = client.post(
"/embed",
json={"text": "Pembayaran berhasil namun e-ticket belum terbit di aplikasi."},
).json()["embedding"]
unrelated = client.post(
"/embed",
json={"text": "Kereta terlambat selama dua jam di stasiun tujuan."},
).json()["embedding"]
assert cosine_similarity(query, similar) > cosine_similarity(query, unrelated)