Spaces:
Sleeping
Sleeping
| 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] | |
| 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) | |