Spaces:
Sleeping
Sleeping
Add TF-IDF ranking and Llama recommendations
Browse files- README.md +7 -4
- app/main.py +31 -11
- app/modules/places/api/dependencies.py +8 -4
- app/modules/places/api/router.py +1 -0
- app/modules/places/api/schemas.py +1 -0
- app/modules/places/application/ports/ranker.py +2 -1
- app/modules/places/application/use_cases/recommend_places.py +48 -4
- app/modules/places/application/use_cases/search_places.py +6 -1
- app/modules/places/domain/models.py +1 -0
- app/modules/places/infrastructure/aws_pgvector_place_repository.py +2 -1
- app/modules/places/infrastructure/mock_place_repository.py +2 -0
- app/modules/places/infrastructure/simple_place_ranker.py +2 -0
- app/modules/places/infrastructure/tfidf_place_ranker.py +119 -0
- app/shared/vector_store/aws_pgvector.py +75 -7
- sql/aws_pgvector_contract.sql +13 -7
- sql/aws_pgvector_full_setup.psql.sql +294 -0
- tests/conftest.py +1 -0
- tests/test_api_endpoints.py +21 -0
- tests/test_places_use_cases.py +74 -4
README.md
CHANGED
|
@@ -7,7 +7,7 @@ pinned: false
|
|
| 7 |
|
| 8 |
# Frimeet API NLP
|
| 9 |
|
| 10 |
-
Servicio NLP independiente para
|
| 11 |
|
| 12 |
La API principal sigue siendo la fuente de verdad de lugares, posts, usuarios, sesiones, permisos y reportes. Este servicio NLP solo trabaja con datos derivados para busqueda semantica.
|
| 13 |
|
|
@@ -22,7 +22,8 @@ Hugging Face API NLP
|
|
| 22 |
|-- usa credenciales nlp_reader
|
| 23 |
|-- consulta RDS PostgreSQL + pgvector
|
| 24 |
|-- genera embedding solo del query del usuario
|
| 25 |
-
|
|
|
|
| 26 |
|
| 27 |
Hugging Face Jobs
|
| 28 |
|-- usan credenciales nlp_writer
|
|
@@ -180,9 +181,11 @@ docs/pgvector_post_embeddings_schema.md
|
|
| 180 |
|
| 181 |
Ese SQL debe ejecutarse una vez con un rol administrador/DBA fuera de Hugging Face. La API NLP usa solo `nlp_reader`; los jobs usan solo `nlp_writer`.
|
| 182 |
|
| 183 |
-
## Llama Via Groq
|
| 184 |
|
| 185 |
-
|
|
|
|
|
|
|
| 186 |
|
| 187 |
El arreglo estructurado `places` viene desde RDS/pgvector mediante embeddings, filtros y ranking. La app debe renderizar cards desde ese arreglo, no parseando texto libre del LLM.
|
| 188 |
|
|
|
|
| 7 |
|
| 8 |
# Frimeet API NLP
|
| 9 |
|
| 10 |
+
Servicio NLP independiente para recuperacion de candidatos con pgvector, ranking TF-IDF, recomendaciones, embeddings y redaccion conversacional con Llama via Groq.
|
| 11 |
|
| 12 |
La API principal sigue siendo la fuente de verdad de lugares, posts, usuarios, sesiones, permisos y reportes. Este servicio NLP solo trabaja con datos derivados para busqueda semantica.
|
| 13 |
|
|
|
|
| 22 |
|-- usa credenciales nlp_reader
|
| 23 |
|-- consulta RDS PostgreSQL + pgvector
|
| 24 |
|-- genera embedding solo del query del usuario
|
| 25 |
+
|-- ordena candidatos con TF-IDF y similitud coseno
|
| 26 |
+
`-- usa Groq/Llama para embellecer recomendaciones y chat
|
| 27 |
|
| 28 |
Hugging Face Jobs
|
| 29 |
|-- usan credenciales nlp_writer
|
|
|
|
| 181 |
|
| 182 |
Ese SQL debe ejecutarse una vez con un rol administrador/DBA fuera de Hugging Face. La API NLP usa solo `nlp_reader`; los jobs usan solo `nlp_writer`.
|
| 183 |
|
| 184 |
+
## Ranking TF-IDF Y Llama Via Groq
|
| 185 |
|
| 186 |
+
`/places/search` y `/places/recommendations` recuperan candidatos filtrados desde pgvector y aplican el flujo TF-IDF de `Lab2_Motor_de_busqueda.ipynb`: TF, IDF, vectorizacion de consulta y similitud coseno. Las etiquetas se ponderan `x6` y la categoria `x2` antes de construir los vectores.
|
| 187 |
+
|
| 188 |
+
Groq/Llama se usa en `/places/recommendations` y `/places/chat` para redactar una respuesta conversacional. No decide que lugares recomendar, no hace busqueda y no inventa lugares.
|
| 189 |
|
| 190 |
El arreglo estructurado `places` viene desde RDS/pgvector mediante embeddings, filtros y ranking. La app debe renderizar cards desde ese arreglo, no parseando texto libre del LLM.
|
| 191 |
|
app/main.py
CHANGED
|
@@ -1,4 +1,5 @@
|
|
| 1 |
from fastapi import FastAPI
|
|
|
|
| 2 |
|
| 3 |
from app.modules.places.api.router import router as places_router
|
| 4 |
from app.modules.posts.api.router import router as posts_router
|
|
@@ -7,6 +8,7 @@ from app.shared.errors.exceptions import AppError
|
|
| 7 |
from app.shared.errors.handlers import app_error_handler
|
| 8 |
from app.shared.logging.config import configure_logging
|
| 9 |
from app.shared.security.request_limits import RequestSizeLimitMiddleware
|
|
|
|
| 10 |
|
| 11 |
|
| 12 |
def create_app() -> FastAPI:
|
|
@@ -38,25 +40,43 @@ def create_app() -> FastAPI:
|
|
| 38 |
async def health() -> dict[str, str]:
|
| 39 |
return {"status": "ok"}
|
| 40 |
|
| 41 |
-
@app.get("/ready", tags=["system"])
|
| 42 |
-
async def ready() -> dict[str, object]:
|
| 43 |
-
|
| 44 |
-
"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 45 |
"environment": settings.env,
|
| 46 |
"dependencies": {
|
| 47 |
-
"vector_store":
|
| 48 |
-
"provider": settings.vector_store_provider,
|
| 49 |
-
"host": settings.pgvector_host,
|
| 50 |
-
"port": settings.pgvector_port,
|
| 51 |
-
"database": settings.pgvector_database,
|
| 52 |
-
"ssl_mode": settings.pgvector_ssl_mode,
|
| 53 |
-
},
|
| 54 |
"llm": {
|
| 55 |
"provider": "groq" if settings.groq_api_key else "mock",
|
| 56 |
"model": settings.groq_model,
|
| 57 |
},
|
| 58 |
},
|
| 59 |
}
|
|
|
|
|
|
|
|
|
|
| 60 |
|
| 61 |
app.include_router(places_router)
|
| 62 |
app.include_router(posts_router)
|
|
|
|
| 1 |
from fastapi import FastAPI
|
| 2 |
+
from fastapi.responses import JSONResponse
|
| 3 |
|
| 4 |
from app.modules.places.api.router import router as places_router
|
| 5 |
from app.modules.posts.api.router import router as posts_router
|
|
|
|
| 8 |
from app.shared.errors.handlers import app_error_handler
|
| 9 |
from app.shared.logging.config import configure_logging
|
| 10 |
from app.shared.security.request_limits import RequestSizeLimitMiddleware
|
| 11 |
+
from app.shared.vector_store.aws_pgvector import AwsPgvectorClient
|
| 12 |
|
| 13 |
|
| 14 |
def create_app() -> FastAPI:
|
|
|
|
| 40 |
async def health() -> dict[str, str]:
|
| 41 |
return {"status": "ok"}
|
| 42 |
|
| 43 |
+
@app.get("/ready", tags=["system"], response_model=None)
|
| 44 |
+
async def ready() -> dict[str, object] | JSONResponse:
|
| 45 |
+
vector_store = {
|
| 46 |
+
"provider": settings.vector_store_provider,
|
| 47 |
+
"host": settings.pgvector_host,
|
| 48 |
+
"port": settings.pgvector_port,
|
| 49 |
+
"database": settings.pgvector_database,
|
| 50 |
+
"ssl_mode": settings.pgvector_ssl_mode,
|
| 51 |
+
}
|
| 52 |
+
is_ready = True
|
| 53 |
+
|
| 54 |
+
if settings.vector_store_provider == "aws_pgvector":
|
| 55 |
+
try:
|
| 56 |
+
contract = await AwsPgvectorClient(settings, role="reader").check_read_contract()
|
| 57 |
+
except Exception as exc:
|
| 58 |
+
contract = {
|
| 59 |
+
"ready": False,
|
| 60 |
+
"error": type(exc).__name__,
|
| 61 |
+
"message": str(exc),
|
| 62 |
+
}
|
| 63 |
+
vector_store["contract"] = contract
|
| 64 |
+
is_ready = bool(contract.get("ready"))
|
| 65 |
+
|
| 66 |
+
payload = {
|
| 67 |
+
"status": "ready" if is_ready else "not_ready",
|
| 68 |
"environment": settings.env,
|
| 69 |
"dependencies": {
|
| 70 |
+
"vector_store": vector_store,
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 71 |
"llm": {
|
| 72 |
"provider": "groq" if settings.groq_api_key else "mock",
|
| 73 |
"model": settings.groq_model,
|
| 74 |
},
|
| 75 |
},
|
| 76 |
}
|
| 77 |
+
if not is_ready:
|
| 78 |
+
return JSONResponse(status_code=503, content=payload)
|
| 79 |
+
return payload
|
| 80 |
|
| 81 |
app.include_router(places_router)
|
| 82 |
app.include_router(posts_router)
|
app/modules/places/api/dependencies.py
CHANGED
|
@@ -7,7 +7,7 @@ from app.modules.places.infrastructure.aws_pgvector_place_repository import (
|
|
| 7 |
AwsPgvectorPlaceRepository,
|
| 8 |
)
|
| 9 |
from app.modules.places.infrastructure.mock_place_repository import MockPlaceVectorRepository
|
| 10 |
-
from app.modules.places.infrastructure.
|
| 11 |
from app.shared.cache.memory import SimpleTTLCache
|
| 12 |
from app.shared.config.settings import get_settings
|
| 13 |
from app.shared.dependencies import get_embedding_provider, get_llm_provider
|
|
@@ -24,8 +24,8 @@ def get_place_repository() -> MockPlaceVectorRepository | AwsPgvectorPlaceReposi
|
|
| 24 |
|
| 25 |
|
| 26 |
@lru_cache
|
| 27 |
-
def get_place_ranker() ->
|
| 28 |
-
return
|
| 29 |
|
| 30 |
|
| 31 |
@lru_cache
|
|
@@ -46,7 +46,11 @@ def get_search_places_use_case() -> SearchPlacesUseCase:
|
|
| 46 |
|
| 47 |
@lru_cache
|
| 48 |
def get_recommend_places_use_case() -> RecommendPlacesUseCase:
|
| 49 |
-
return RecommendPlacesUseCase(
|
|
|
|
|
|
|
|
|
|
|
|
|
| 50 |
|
| 51 |
|
| 52 |
@lru_cache
|
|
|
|
| 7 |
AwsPgvectorPlaceRepository,
|
| 8 |
)
|
| 9 |
from app.modules.places.infrastructure.mock_place_repository import MockPlaceVectorRepository
|
| 10 |
+
from app.modules.places.infrastructure.tfidf_place_ranker import TfidfPlaceRanker
|
| 11 |
from app.shared.cache.memory import SimpleTTLCache
|
| 12 |
from app.shared.config.settings import get_settings
|
| 13 |
from app.shared.dependencies import get_embedding_provider, get_llm_provider
|
|
|
|
| 24 |
|
| 25 |
|
| 26 |
@lru_cache
|
| 27 |
+
def get_place_ranker() -> TfidfPlaceRanker:
|
| 28 |
+
return TfidfPlaceRanker()
|
| 29 |
|
| 30 |
|
| 31 |
@lru_cache
|
|
|
|
| 46 |
|
| 47 |
@lru_cache
|
| 48 |
def get_recommend_places_use_case() -> RecommendPlacesUseCase:
|
| 49 |
+
return RecommendPlacesUseCase(
|
| 50 |
+
search_use_case=get_search_places_use_case(),
|
| 51 |
+
llm_provider=get_llm_provider(),
|
| 52 |
+
output_guard=PlaceChatOutputGuard(),
|
| 53 |
+
)
|
| 54 |
|
| 55 |
|
| 56 |
@lru_cache
|
app/modules/places/api/router.py
CHANGED
|
@@ -54,6 +54,7 @@ async def recommend_places(
|
|
| 54 |
)
|
| 55 |
return PlaceRecommendationResponse(
|
| 56 |
query=result.query,
|
|
|
|
| 57 |
places=[place_to_schema(place) for place in result.places],
|
| 58 |
metadata=result.metadata,
|
| 59 |
)
|
|
|
|
| 54 |
)
|
| 55 |
return PlaceRecommendationResponse(
|
| 56 |
query=result.query,
|
| 57 |
+
message=result.message,
|
| 58 |
places=[place_to_schema(place) for place in result.places],
|
| 59 |
metadata=result.metadata,
|
| 60 |
)
|
app/modules/places/api/schemas.py
CHANGED
|
@@ -77,6 +77,7 @@ class PlaceSearchResponse(BaseModel):
|
|
| 77 |
|
| 78 |
class PlaceRecommendationResponse(BaseModel):
|
| 79 |
query: str
|
|
|
|
| 80 |
places: list[PlaceResultSchema]
|
| 81 |
metadata: dict[str, Any] = Field(default_factory=dict)
|
| 82 |
|
|
|
|
| 77 |
|
| 78 |
class PlaceRecommendationResponse(BaseModel):
|
| 79 |
query: str
|
| 80 |
+
message: str
|
| 81 |
places: list[PlaceResultSchema]
|
| 82 |
metadata: dict[str, Any] = Field(default_factory=dict)
|
| 83 |
|
app/modules/places/application/ports/ranker.py
CHANGED
|
@@ -6,8 +6,9 @@ from app.modules.places.domain.models import PlaceCandidate, PlaceFilters
|
|
| 6 |
class PlaceRanker(Protocol):
|
| 7 |
def rank(
|
| 8 |
self,
|
|
|
|
| 9 |
places: Sequence[PlaceCandidate],
|
| 10 |
filters: PlaceFilters,
|
| 11 |
limit: int,
|
| 12 |
) -> list[PlaceCandidate]:
|
| 13 |
-
"""Rank candidates
|
|
|
|
| 6 |
class PlaceRanker(Protocol):
|
| 7 |
def rank(
|
| 8 |
self,
|
| 9 |
+
query: str,
|
| 10 |
places: Sequence[PlaceCandidate],
|
| 11 |
filters: PlaceFilters,
|
| 12 |
limit: int,
|
| 13 |
) -> list[PlaceCandidate]:
|
| 14 |
+
"""Rank place candidates for a normalized user query."""
|
app/modules/places/application/use_cases/recommend_places.py
CHANGED
|
@@ -1,20 +1,31 @@
|
|
| 1 |
from dataclasses import dataclass, field
|
|
|
|
| 2 |
from typing import Any
|
| 3 |
|
| 4 |
from app.modules.places.application.use_cases.search_places import SearchPlacesUseCase
|
| 5 |
from app.modules.places.domain.models import PlaceCandidate, PlaceFilters
|
|
|
|
|
|
|
| 6 |
|
| 7 |
|
| 8 |
@dataclass(frozen=True)
|
| 9 |
class RecommendPlacesResult:
|
| 10 |
query: str
|
|
|
|
| 11 |
places: list[PlaceCandidate]
|
| 12 |
metadata: dict[str, Any] = field(default_factory=dict)
|
| 13 |
|
| 14 |
|
| 15 |
class RecommendPlacesUseCase:
|
| 16 |
-
def __init__(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 17 |
self._search_use_case = search_use_case
|
|
|
|
|
|
|
| 18 |
|
| 19 |
async def execute(
|
| 20 |
self,
|
|
@@ -27,11 +38,44 @@ class RecommendPlacesUseCase:
|
|
| 27 |
filters=filters,
|
| 28 |
limit=limit,
|
| 29 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 30 |
return RecommendPlacesResult(
|
| 31 |
query=search_result.query,
|
| 32 |
-
|
|
|
|
| 33 |
metadata={
|
| 34 |
-
"strategy": "
|
| 35 |
-
"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 36 |
},
|
| 37 |
)
|
|
|
|
| 1 |
from dataclasses import dataclass, field
|
| 2 |
+
from datetime import UTC, datetime
|
| 3 |
from typing import Any
|
| 4 |
|
| 5 |
from app.modules.places.application.use_cases.search_places import SearchPlacesUseCase
|
| 6 |
from app.modules.places.domain.models import PlaceCandidate, PlaceFilters
|
| 7 |
+
from app.shared.nlp.llm.base import LLMProvider
|
| 8 |
+
from app.shared.nlp.llm.output_guard import PlaceChatOutputGuard
|
| 9 |
|
| 10 |
|
| 11 |
@dataclass(frozen=True)
|
| 12 |
class RecommendPlacesResult:
|
| 13 |
query: str
|
| 14 |
+
message: str
|
| 15 |
places: list[PlaceCandidate]
|
| 16 |
metadata: dict[str, Any] = field(default_factory=dict)
|
| 17 |
|
| 18 |
|
| 19 |
class RecommendPlacesUseCase:
|
| 20 |
+
def __init__(
|
| 21 |
+
self,
|
| 22 |
+
search_use_case: SearchPlacesUseCase,
|
| 23 |
+
llm_provider: LLMProvider,
|
| 24 |
+
output_guard: PlaceChatOutputGuard,
|
| 25 |
+
) -> None:
|
| 26 |
self._search_use_case = search_use_case
|
| 27 |
+
self._llm_provider = llm_provider
|
| 28 |
+
self._output_guard = output_guard
|
| 29 |
|
| 30 |
async def execute(
|
| 31 |
self,
|
|
|
|
| 38 |
filters=filters,
|
| 39 |
limit=limit,
|
| 40 |
)
|
| 41 |
+
places = search_result.places
|
| 42 |
+
llm_provider = self._llm_provider.provider_name
|
| 43 |
+
llm_model = self._llm_provider.model_name
|
| 44 |
+
used_llm = False
|
| 45 |
+
guard_reason = None
|
| 46 |
+
|
| 47 |
+
try:
|
| 48 |
+
llm_result = await self._llm_provider.generate_place_chat_response(
|
| 49 |
+
user_intent=search_result.normalized_query,
|
| 50 |
+
region=filters.city or filters.state,
|
| 51 |
+
places=[place.to_llm_context() for place in places],
|
| 52 |
+
)
|
| 53 |
+
llm_provider = llm_result.provider
|
| 54 |
+
llm_model = llm_result.model
|
| 55 |
+
guarded = self._output_guard.validate(
|
| 56 |
+
message=llm_result.message,
|
| 57 |
+
allowed_place_names=[place.name for place in places],
|
| 58 |
+
)
|
| 59 |
+
message = guarded.message
|
| 60 |
+
used_llm = not guarded.used_fallback
|
| 61 |
+
guard_reason = guarded.reason
|
| 62 |
+
except Exception as exc:
|
| 63 |
+
guarded = self._output_guard.fallback(reason=exc.__class__.__name__)
|
| 64 |
+
message = guarded.message
|
| 65 |
+
guard_reason = guarded.reason
|
| 66 |
+
|
| 67 |
return RecommendPlacesResult(
|
| 68 |
query=search_result.query,
|
| 69 |
+
message=message,
|
| 70 |
+
places=places,
|
| 71 |
metadata={
|
| 72 |
+
"strategy": "pgvector_candidates_plus_tfidf_ranking",
|
| 73 |
+
"ranking": "tfidf_cosine",
|
| 74 |
+
"llm_provider": llm_provider,
|
| 75 |
+
"llm_model": llm_model,
|
| 76 |
+
"used_llm": used_llm,
|
| 77 |
+
"guard_reason": guard_reason,
|
| 78 |
+
"places_used_as_context": [place.id for place in places],
|
| 79 |
+
"timestamp": datetime.now(UTC).isoformat(),
|
| 80 |
},
|
| 81 |
)
|
app/modules/places/application/use_cases/search_places.py
CHANGED
|
@@ -49,7 +49,12 @@ class SearchPlacesUseCase:
|
|
| 49 |
filters=filters,
|
| 50 |
limit=max(limit * 3, limit),
|
| 51 |
)
|
| 52 |
-
ranked_places = self._ranker.rank(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 53 |
result = SearchPlacesResult(
|
| 54 |
query=query,
|
| 55 |
normalized_query=normalized_query,
|
|
|
|
| 49 |
filters=filters,
|
| 50 |
limit=max(limit * 3, limit),
|
| 51 |
)
|
| 52 |
+
ranked_places = self._ranker.rank(
|
| 53 |
+
query=normalized_query,
|
| 54 |
+
places=candidates,
|
| 55 |
+
filters=filters,
|
| 56 |
+
limit=limit,
|
| 57 |
+
)
|
| 58 |
result = SearchPlacesResult(
|
| 59 |
query=query,
|
| 60 |
normalized_query=normalized_query,
|
app/modules/places/domain/models.py
CHANGED
|
@@ -36,6 +36,7 @@ class PlaceCandidate:
|
|
| 36 |
state: str | None = None
|
| 37 |
price_range: str | None = None
|
| 38 |
metadata: dict[str, Any] = field(default_factory=dict)
|
|
|
|
| 39 |
|
| 40 |
def to_llm_context(self) -> dict[str, Any]:
|
| 41 |
context = {
|
|
|
|
| 36 |
state: str | None = None
|
| 37 |
price_range: str | None = None
|
| 38 |
metadata: dict[str, Any] = field(default_factory=dict)
|
| 39 |
+
document: str | None = None
|
| 40 |
|
| 41 |
def to_llm_context(self) -> dict[str, Any]:
|
| 42 |
context = {
|
app/modules/places/infrastructure/aws_pgvector_place_repository.py
CHANGED
|
@@ -27,7 +27,7 @@ class AwsPgvectorPlaceRepository(PlaceVectorRepository):
|
|
| 27 |
|
| 28 |
|
| 29 |
def _match_to_candidate(match: VectorMatch) -> PlaceCandidate:
|
| 30 |
-
metadata = match.metadata
|
| 31 |
return PlaceCandidate(
|
| 32 |
id=match.id,
|
| 33 |
name=str(metadata.get("name") or match.id),
|
|
@@ -37,4 +37,5 @@ def _match_to_candidate(match: VectorMatch) -> PlaceCandidate:
|
|
| 37 |
state=metadata.get("state"),
|
| 38 |
price_range=metadata.get("price_range"),
|
| 39 |
metadata=metadata,
|
|
|
|
| 40 |
)
|
|
|
|
| 27 |
|
| 28 |
|
| 29 |
def _match_to_candidate(match: VectorMatch) -> PlaceCandidate:
|
| 30 |
+
metadata = dict(match.metadata)
|
| 31 |
return PlaceCandidate(
|
| 32 |
id=match.id,
|
| 33 |
name=str(metadata.get("name") or match.id),
|
|
|
|
| 37 |
state=metadata.get("state"),
|
| 38 |
price_range=metadata.get("price_range"),
|
| 39 |
metadata=metadata,
|
| 40 |
+
document=match.document,
|
| 41 |
)
|
app/modules/places/infrastructure/mock_place_repository.py
CHANGED
|
@@ -78,6 +78,7 @@ class MockPlaceVectorRepository(PlaceVectorRepository):
|
|
| 78 |
self._records.append(
|
| 79 |
{
|
| 80 |
"place": place,
|
|
|
|
| 81 |
"embedding": embedding_provider.embed_text(
|
| 82 |
prepare_for_embedding(searchable_text)
|
| 83 |
),
|
|
@@ -111,6 +112,7 @@ class MockPlaceVectorRepository(PlaceVectorRepository):
|
|
| 111 |
"tags": place["tags"],
|
| 112 |
"short_description": place["short_description"],
|
| 113 |
},
|
|
|
|
| 114 |
)
|
| 115 |
)
|
| 116 |
return sorted(candidates, key=lambda item: item.score, reverse=True)[:limit]
|
|
|
|
| 78 |
self._records.append(
|
| 79 |
{
|
| 80 |
"place": place,
|
| 81 |
+
"document": searchable_text,
|
| 82 |
"embedding": embedding_provider.embed_text(
|
| 83 |
prepare_for_embedding(searchable_text)
|
| 84 |
),
|
|
|
|
| 112 |
"tags": place["tags"],
|
| 113 |
"short_description": place["short_description"],
|
| 114 |
},
|
| 115 |
+
document=record["document"],
|
| 116 |
)
|
| 117 |
)
|
| 118 |
return sorted(candidates, key=lambda item: item.score, reverse=True)[:limit]
|
app/modules/places/infrastructure/simple_place_ranker.py
CHANGED
|
@@ -7,10 +7,12 @@ from app.modules.places.domain.models import PlaceCandidate, PlaceFilters
|
|
| 7 |
class SimplePlaceRanker(PlaceRanker):
|
| 8 |
def rank(
|
| 9 |
self,
|
|
|
|
| 10 |
places: Sequence[PlaceCandidate],
|
| 11 |
filters: PlaceFilters,
|
| 12 |
limit: int,
|
| 13 |
) -> list[PlaceCandidate]:
|
|
|
|
| 14 |
ranked = sorted(
|
| 15 |
places,
|
| 16 |
key=lambda place: (
|
|
|
|
| 7 |
class SimplePlaceRanker(PlaceRanker):
|
| 8 |
def rank(
|
| 9 |
self,
|
| 10 |
+
query: str,
|
| 11 |
places: Sequence[PlaceCandidate],
|
| 12 |
filters: PlaceFilters,
|
| 13 |
limit: int,
|
| 14 |
) -> list[PlaceCandidate]:
|
| 15 |
+
del query
|
| 16 |
ranked = sorted(
|
| 17 |
places,
|
| 18 |
key=lambda place: (
|
app/modules/places/infrastructure/tfidf_place_ranker.py
ADDED
|
@@ -0,0 +1,119 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from collections import Counter
|
| 2 |
+
from dataclasses import replace
|
| 3 |
+
import math
|
| 4 |
+
import re
|
| 5 |
+
from typing import Any, Sequence
|
| 6 |
+
|
| 7 |
+
from app.modules.places.application.ports.ranker import PlaceRanker
|
| 8 |
+
from app.modules.places.domain.models import PlaceCandidate, PlaceFilters
|
| 9 |
+
from app.shared.nlp.preprocessing.text import prepare_for_embedding
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
TAG_WEIGHT = 6
|
| 13 |
+
CATEGORY_WEIGHT = 2
|
| 14 |
+
TOKEN_PATTERN = re.compile(r"[a-z0-9]+")
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
class TfidfPlaceRanker(PlaceRanker):
|
| 18 |
+
"""Rerank retrieved candidates with the TF-IDF flow from Lab 2."""
|
| 19 |
+
|
| 20 |
+
def rank(
|
| 21 |
+
self,
|
| 22 |
+
query: str,
|
| 23 |
+
places: Sequence[PlaceCandidate],
|
| 24 |
+
filters: PlaceFilters,
|
| 25 |
+
limit: int,
|
| 26 |
+
) -> list[PlaceCandidate]:
|
| 27 |
+
del filters
|
| 28 |
+
candidates = list(places)
|
| 29 |
+
if not candidates:
|
| 30 |
+
return []
|
| 31 |
+
|
| 32 |
+
corpus = [_place_tokens(place) for place in candidates]
|
| 33 |
+
idf_index = idf(corpus)
|
| 34 |
+
query_vector = tfidf(_tokenize(query), idf_index)
|
| 35 |
+
|
| 36 |
+
scored = [
|
| 37 |
+
(cosine_similarity(query_vector, tfidf(document, idf_index)), place)
|
| 38 |
+
for place, document in zip(candidates, corpus)
|
| 39 |
+
]
|
| 40 |
+
scored.sort(key=lambda item: (item[0], item[1].score), reverse=True)
|
| 41 |
+
|
| 42 |
+
return [
|
| 43 |
+
replace(place, score=score)
|
| 44 |
+
for score, place in scored[:limit]
|
| 45 |
+
]
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
def tf(document: Sequence[str]) -> dict[str, float]:
|
| 49 |
+
total = len(document)
|
| 50 |
+
if total == 0:
|
| 51 |
+
return {}
|
| 52 |
+
counts = Counter(document)
|
| 53 |
+
return {term: count / total for term, count in counts.items()}
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
def idf(corpus: Sequence[Sequence[str]]) -> dict[str, float]:
|
| 57 |
+
document_count = len(corpus)
|
| 58 |
+
if document_count == 0:
|
| 59 |
+
return {}
|
| 60 |
+
|
| 61 |
+
document_frequency: Counter[str] = Counter()
|
| 62 |
+
for document in corpus:
|
| 63 |
+
document_frequency.update(set(document))
|
| 64 |
+
return {
|
| 65 |
+
term: math.log(document_count / frequency)
|
| 66 |
+
for term, frequency in document_frequency.items()
|
| 67 |
+
}
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
def tfidf(document: Sequence[str], idf_index: dict[str, float]) -> dict[str, float]:
|
| 71 |
+
return {
|
| 72 |
+
term: frequency * idf_index.get(term, 0.0)
|
| 73 |
+
for term, frequency in tf(document).items()
|
| 74 |
+
}
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
def cosine_similarity(left: dict[str, float], right: dict[str, float]) -> float:
|
| 78 |
+
dot_product = sum(weight * right.get(term, 0.0) for term, weight in left.items())
|
| 79 |
+
left_norm = math.sqrt(sum(weight**2 for weight in left.values()))
|
| 80 |
+
right_norm = math.sqrt(sum(weight**2 for weight in right.values()))
|
| 81 |
+
if left_norm == 0 or right_norm == 0:
|
| 82 |
+
return 0.0
|
| 83 |
+
return dot_product / (left_norm * right_norm)
|
| 84 |
+
|
| 85 |
+
|
| 86 |
+
def _place_tokens(place: PlaceCandidate) -> list[str]:
|
| 87 |
+
tags = _as_text(place.metadata.get("tags"))
|
| 88 |
+
category = place.category or ""
|
| 89 |
+
base_document = place.document or " ".join(
|
| 90 |
+
value
|
| 91 |
+
for value in [
|
| 92 |
+
place.name,
|
| 93 |
+
category,
|
| 94 |
+
place.city or "",
|
| 95 |
+
place.state or "",
|
| 96 |
+
tags,
|
| 97 |
+
_as_text(place.metadata.get("occasion")),
|
| 98 |
+
_as_text(place.metadata.get("short_description")),
|
| 99 |
+
]
|
| 100 |
+
if value
|
| 101 |
+
)
|
| 102 |
+
|
| 103 |
+
weighted_fields = [base_document]
|
| 104 |
+
weighted_fields.extend([category] * (CATEGORY_WEIGHT - 1))
|
| 105 |
+
weighted_fields.extend([tags] * (TAG_WEIGHT - 1))
|
| 106 |
+
return _tokenize(" ".join(value for value in weighted_fields if value))
|
| 107 |
+
|
| 108 |
+
|
| 109 |
+
def _tokenize(text: str) -> list[str]:
|
| 110 |
+
normalized = prepare_for_embedding(text)
|
| 111 |
+
return TOKEN_PATTERN.findall(normalized)
|
| 112 |
+
|
| 113 |
+
|
| 114 |
+
def _as_text(value: Any) -> str:
|
| 115 |
+
if value is None:
|
| 116 |
+
return ""
|
| 117 |
+
if isinstance(value, (list, tuple, set)):
|
| 118 |
+
return " ".join(str(item) for item in value)
|
| 119 |
+
return str(value)
|
app/shared/vector_store/aws_pgvector.py
CHANGED
|
@@ -7,10 +7,17 @@ from typing import Any, AsyncIterator
|
|
| 7 |
import asyncpg
|
| 8 |
|
| 9 |
from app.shared.config.settings import Settings
|
|
|
|
| 10 |
from app.shared.vector_store.models import VectorMatch, VectorUpsertRecord
|
| 11 |
from app.shared.vector_store.sql import quote_identifier, vector_literal
|
| 12 |
|
| 13 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 14 |
class AwsPgvectorClient:
|
| 15 |
"""PostgreSQL + pgvector access for RDS/Aurora using controlled SQL functions."""
|
| 16 |
|
|
@@ -54,6 +61,58 @@ class AwsPgvectorClient:
|
|
| 54 |
limit=limit,
|
| 55 |
)
|
| 56 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 57 |
async def fetch_place_content_hashes(self, ids: Iterable[str]) -> dict[str, str]:
|
| 58 |
return await self._fetch_content_hashes(
|
| 59 |
function_name="get_place_content_hashes",
|
|
@@ -92,13 +151,22 @@ class AwsPgvectorClient:
|
|
| 92 |
limit: int,
|
| 93 |
) -> list[VectorMatch]:
|
| 94 |
query = f"SELECT * FROM {quote_identifier(function_name)}($1::vector, $2::integer, $3::jsonb)"
|
| 95 |
-
|
| 96 |
-
|
| 97 |
-
|
| 98 |
-
|
| 99 |
-
|
| 100 |
-
|
| 101 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 102 |
return [_row_to_vector_match(row) for row in rows]
|
| 103 |
|
| 104 |
async def _fetch_content_hashes(
|
|
|
|
| 7 |
import asyncpg
|
| 8 |
|
| 9 |
from app.shared.config.settings import Settings
|
| 10 |
+
from app.shared.errors.exceptions import AppError
|
| 11 |
from app.shared.vector_store.models import VectorMatch, VectorUpsertRecord
|
| 12 |
from app.shared.vector_store.sql import quote_identifier, vector_literal
|
| 13 |
|
| 14 |
|
| 15 |
+
READ_CONTRACT_SIGNATURES = {
|
| 16 |
+
"match_places": "match_places(vector, integer, jsonb)",
|
| 17 |
+
"match_posts": "match_posts(vector, integer, jsonb)",
|
| 18 |
+
}
|
| 19 |
+
|
| 20 |
+
|
| 21 |
class AwsPgvectorClient:
|
| 22 |
"""PostgreSQL + pgvector access for RDS/Aurora using controlled SQL functions."""
|
| 23 |
|
|
|
|
| 61 |
limit=limit,
|
| 62 |
)
|
| 63 |
|
| 64 |
+
async def check_read_contract(self) -> dict[str, Any]:
|
| 65 |
+
"""Check that read-only pgvector functions are visible and executable."""
|
| 66 |
+
try:
|
| 67 |
+
async with self.connection() as connection:
|
| 68 |
+
vector_row = await connection.fetchrow(
|
| 69 |
+
"SELECT to_regtype('vector') IS NOT NULL AS available"
|
| 70 |
+
)
|
| 71 |
+
vector_available = bool(vector_row["available"]) if vector_row else False
|
| 72 |
+
functions: dict[str, dict[str, Any]] = {
|
| 73 |
+
function_name: {
|
| 74 |
+
"signature": signature,
|
| 75 |
+
"exists": False,
|
| 76 |
+
"executable": False,
|
| 77 |
+
}
|
| 78 |
+
for function_name, signature in READ_CONTRACT_SIGNATURES.items()
|
| 79 |
+
}
|
| 80 |
+
if vector_available:
|
| 81 |
+
for function_name, signature in READ_CONTRACT_SIGNATURES.items():
|
| 82 |
+
row = await connection.fetchrow(
|
| 83 |
+
"""
|
| 84 |
+
SELECT
|
| 85 |
+
to_regprocedure($1) IS NOT NULL AS exists,
|
| 86 |
+
COALESCE(
|
| 87 |
+
has_function_privilege(to_regprocedure($1), 'EXECUTE'),
|
| 88 |
+
false
|
| 89 |
+
) AS executable
|
| 90 |
+
""",
|
| 91 |
+
signature,
|
| 92 |
+
)
|
| 93 |
+
functions[function_name].update(
|
| 94 |
+
{
|
| 95 |
+
"exists": bool(row["exists"]) if row else False,
|
| 96 |
+
"executable": bool(row["executable"]) if row else False,
|
| 97 |
+
}
|
| 98 |
+
)
|
| 99 |
+
except Exception as exc:
|
| 100 |
+
return {
|
| 101 |
+
"ready": False,
|
| 102 |
+
"error": type(exc).__name__,
|
| 103 |
+
"message": str(exc),
|
| 104 |
+
}
|
| 105 |
+
|
| 106 |
+
ready = vector_available and all(
|
| 107 |
+
details["exists"] and details["executable"]
|
| 108 |
+
for details in functions.values()
|
| 109 |
+
)
|
| 110 |
+
return {
|
| 111 |
+
"ready": ready,
|
| 112 |
+
"vector_extension": vector_available,
|
| 113 |
+
"functions": functions,
|
| 114 |
+
}
|
| 115 |
+
|
| 116 |
async def fetch_place_content_hashes(self, ids: Iterable[str]) -> dict[str, str]:
|
| 117 |
return await self._fetch_content_hashes(
|
| 118 |
function_name="get_place_content_hashes",
|
|
|
|
| 151 |
limit: int,
|
| 152 |
) -> list[VectorMatch]:
|
| 153 |
query = f"SELECT * FROM {quote_identifier(function_name)}($1::vector, $2::integer, $3::jsonb)"
|
| 154 |
+
try:
|
| 155 |
+
async with self.connection() as connection:
|
| 156 |
+
rows = await connection.fetch(
|
| 157 |
+
query,
|
| 158 |
+
vector_literal(embedding),
|
| 159 |
+
limit,
|
| 160 |
+
json.dumps(filters, ensure_ascii=False),
|
| 161 |
+
)
|
| 162 |
+
except asyncpg.exceptions.UndefinedFunctionError as exc:
|
| 163 |
+
raise AppError(
|
| 164 |
+
"Pgvector SQL contract is missing or not visible to this role. "
|
| 165 |
+
"Run sql/aws_pgvector_contract.sql in the configured database and "
|
| 166 |
+
"grant EXECUTE to the reader role.",
|
| 167 |
+
code="pgvector_contract_missing",
|
| 168 |
+
status_code=503,
|
| 169 |
+
) from exc
|
| 170 |
return [_row_to_vector_match(row) for row in rows]
|
| 171 |
|
| 172 |
async def _fetch_content_hashes(
|
sql/aws_pgvector_contract.sql
CHANGED
|
@@ -53,20 +53,21 @@ RETURNS TABLE (
|
|
| 53 |
LANGUAGE sql
|
| 54 |
STABLE
|
| 55 |
SECURITY DEFINER
|
|
|
|
| 56 |
AS $$
|
| 57 |
SELECT
|
| 58 |
p.external_id,
|
| 59 |
p.document,
|
| 60 |
p.metadata,
|
| 61 |
1 - (p.embedding <=> query_embedding) AS score
|
| 62 |
-
FROM place_embeddings p
|
| 63 |
WHERE p.is_active = true
|
| 64 |
AND COALESCE((filters->>'is_active')::boolean, true) = true
|
| 65 |
AND ((filters ? 'city') IS FALSE OR lower(p.metadata->>'city') = lower(filters->>'city'))
|
| 66 |
AND ((filters ? 'state') IS FALSE OR lower(p.metadata->>'state') = lower(filters->>'state'))
|
| 67 |
AND ((filters ? 'category') IS FALSE OR lower(p.metadata->>'category') = lower(filters->>'category'))
|
| 68 |
AND ((filters ? 'price_range') IS FALSE OR p.metadata->>'price_range' = filters->>'price_range')
|
| 69 |
-
AND ((filters ? 'occasion') IS FALSE OR p.metadata->>'occasion' ILIKE '%' || filters->>'occasion' || '%')
|
| 70 |
ORDER BY p.embedding <=> query_embedding
|
| 71 |
LIMIT match_count;
|
| 72 |
$$;
|
|
@@ -85,13 +86,14 @@ RETURNS TABLE (
|
|
| 85 |
LANGUAGE sql
|
| 86 |
STABLE
|
| 87 |
SECURITY DEFINER
|
|
|
|
| 88 |
AS $$
|
| 89 |
SELECT
|
| 90 |
p.external_id,
|
| 91 |
p.document,
|
| 92 |
p.metadata,
|
| 93 |
1 - (p.embedding <=> query_embedding) AS score
|
| 94 |
-
FROM post_embeddings p
|
| 95 |
WHERE p.is_active = true
|
| 96 |
AND COALESCE((filters->>'is_active')::boolean, true) = true
|
| 97 |
AND ((filters ? 'city') IS FALSE OR lower(p.metadata->>'city') = lower(filters->>'city'))
|
|
@@ -112,8 +114,9 @@ CREATE OR REPLACE FUNCTION upsert_place_embedding(
|
|
| 112 |
RETURNS VOID
|
| 113 |
LANGUAGE sql
|
| 114 |
SECURITY DEFINER
|
|
|
|
| 115 |
AS $$
|
| 116 |
-
INSERT INTO place_embeddings (
|
| 117 |
external_id,
|
| 118 |
document,
|
| 119 |
metadata,
|
|
@@ -159,8 +162,9 @@ CREATE OR REPLACE FUNCTION upsert_post_embedding(
|
|
| 159 |
RETURNS VOID
|
| 160 |
LANGUAGE sql
|
| 161 |
SECURITY DEFINER
|
|
|
|
| 162 |
AS $$
|
| 163 |
-
INSERT INTO post_embeddings (
|
| 164 |
external_id,
|
| 165 |
document,
|
| 166 |
metadata,
|
|
@@ -201,9 +205,10 @@ RETURNS TABLE (
|
|
| 201 |
LANGUAGE sql
|
| 202 |
STABLE
|
| 203 |
SECURITY DEFINER
|
|
|
|
| 204 |
AS $$
|
| 205 |
SELECT p.external_id, p.content_hash
|
| 206 |
-
FROM place_embeddings p
|
| 207 |
WHERE p.external_id = ANY(p_external_ids);
|
| 208 |
$$;
|
| 209 |
|
|
@@ -215,9 +220,10 @@ RETURNS TABLE (
|
|
| 215 |
LANGUAGE sql
|
| 216 |
STABLE
|
| 217 |
SECURITY DEFINER
|
|
|
|
| 218 |
AS $$
|
| 219 |
SELECT p.external_id, p.content_hash
|
| 220 |
-
FROM post_embeddings p
|
| 221 |
WHERE p.external_id = ANY(p_external_ids);
|
| 222 |
$$;
|
| 223 |
|
|
|
|
| 53 |
LANGUAGE sql
|
| 54 |
STABLE
|
| 55 |
SECURITY DEFINER
|
| 56 |
+
SET search_path = public
|
| 57 |
AS $$
|
| 58 |
SELECT
|
| 59 |
p.external_id,
|
| 60 |
p.document,
|
| 61 |
p.metadata,
|
| 62 |
1 - (p.embedding <=> query_embedding) AS score
|
| 63 |
+
FROM public.place_embeddings p
|
| 64 |
WHERE p.is_active = true
|
| 65 |
AND COALESCE((filters->>'is_active')::boolean, true) = true
|
| 66 |
AND ((filters ? 'city') IS FALSE OR lower(p.metadata->>'city') = lower(filters->>'city'))
|
| 67 |
AND ((filters ? 'state') IS FALSE OR lower(p.metadata->>'state') = lower(filters->>'state'))
|
| 68 |
AND ((filters ? 'category') IS FALSE OR lower(p.metadata->>'category') = lower(filters->>'category'))
|
| 69 |
AND ((filters ? 'price_range') IS FALSE OR p.metadata->>'price_range' = filters->>'price_range')
|
| 70 |
+
AND ((filters ? 'occasion') IS FALSE OR p.metadata->>'occasion' ILIKE ('%' || (filters->>'occasion') || '%'))
|
| 71 |
ORDER BY p.embedding <=> query_embedding
|
| 72 |
LIMIT match_count;
|
| 73 |
$$;
|
|
|
|
| 86 |
LANGUAGE sql
|
| 87 |
STABLE
|
| 88 |
SECURITY DEFINER
|
| 89 |
+
SET search_path = public
|
| 90 |
AS $$
|
| 91 |
SELECT
|
| 92 |
p.external_id,
|
| 93 |
p.document,
|
| 94 |
p.metadata,
|
| 95 |
1 - (p.embedding <=> query_embedding) AS score
|
| 96 |
+
FROM public.post_embeddings p
|
| 97 |
WHERE p.is_active = true
|
| 98 |
AND COALESCE((filters->>'is_active')::boolean, true) = true
|
| 99 |
AND ((filters ? 'city') IS FALSE OR lower(p.metadata->>'city') = lower(filters->>'city'))
|
|
|
|
| 114 |
RETURNS VOID
|
| 115 |
LANGUAGE sql
|
| 116 |
SECURITY DEFINER
|
| 117 |
+
SET search_path = public
|
| 118 |
AS $$
|
| 119 |
+
INSERT INTO public.place_embeddings (
|
| 120 |
external_id,
|
| 121 |
document,
|
| 122 |
metadata,
|
|
|
|
| 162 |
RETURNS VOID
|
| 163 |
LANGUAGE sql
|
| 164 |
SECURITY DEFINER
|
| 165 |
+
SET search_path = public
|
| 166 |
AS $$
|
| 167 |
+
INSERT INTO public.post_embeddings (
|
| 168 |
external_id,
|
| 169 |
document,
|
| 170 |
metadata,
|
|
|
|
| 205 |
LANGUAGE sql
|
| 206 |
STABLE
|
| 207 |
SECURITY DEFINER
|
| 208 |
+
SET search_path = public
|
| 209 |
AS $$
|
| 210 |
SELECT p.external_id, p.content_hash
|
| 211 |
+
FROM public.place_embeddings p
|
| 212 |
WHERE p.external_id = ANY(p_external_ids);
|
| 213 |
$$;
|
| 214 |
|
|
|
|
| 220 |
LANGUAGE sql
|
| 221 |
STABLE
|
| 222 |
SECURITY DEFINER
|
| 223 |
+
SET search_path = public
|
| 224 |
AS $$
|
| 225 |
SELECT p.external_id, p.content_hash
|
| 226 |
+
FROM public.post_embeddings p
|
| 227 |
WHERE p.external_id = ANY(p_external_ids);
|
| 228 |
$$;
|
| 229 |
|
sql/aws_pgvector_full_setup.psql.sql
ADDED
|
@@ -0,0 +1,294 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
-- Full setup for Frimeet API NLP pgvector storage.
|
| 2 |
+
-- Run with psql using an admin/RDS master role:
|
| 3 |
+
-- psql "host=<host> port=5432 dbname=postgres user=<admin> sslmode=require" -f sql/aws_pgvector_full_setup.psql.sql
|
| 4 |
+
--
|
| 5 |
+
-- Replace the passwords before running.
|
| 6 |
+
-- VECTOR(16) must match EMBEDDING_DIMENSION=16 in the API environment.
|
| 7 |
+
|
| 8 |
+
\set ON_ERROR_STOP on
|
| 9 |
+
|
| 10 |
+
DO $$
|
| 11 |
+
BEGIN
|
| 12 |
+
IF NOT EXISTS (SELECT 1 FROM pg_roles WHERE rolname = 'nlp_owner') THEN
|
| 13 |
+
CREATE ROLE nlp_owner LOGIN PASSWORD 'CAMBIA_OWNER_PASSWORD';
|
| 14 |
+
END IF;
|
| 15 |
+
|
| 16 |
+
IF NOT EXISTS (SELECT 1 FROM pg_roles WHERE rolname = 'nlp_reader') THEN
|
| 17 |
+
CREATE ROLE nlp_reader LOGIN PASSWORD 'CAMBIA_READER_PASSWORD';
|
| 18 |
+
END IF;
|
| 19 |
+
|
| 20 |
+
IF NOT EXISTS (SELECT 1 FROM pg_roles WHERE rolname = 'nlp_writer') THEN
|
| 21 |
+
CREATE ROLE nlp_writer LOGIN PASSWORD 'CAMBIA_WRITER_PASSWORD';
|
| 22 |
+
END IF;
|
| 23 |
+
END $$;
|
| 24 |
+
|
| 25 |
+
SELECT 'CREATE DATABASE nlp_vectors OWNER nlp_owner'
|
| 26 |
+
WHERE NOT EXISTS (SELECT 1 FROM pg_database WHERE datname = 'nlp_vectors')
|
| 27 |
+
\gexec
|
| 28 |
+
|
| 29 |
+
GRANT CONNECT ON DATABASE nlp_vectors TO nlp_reader, nlp_writer;
|
| 30 |
+
ALTER ROLE nlp_reader SET search_path = public;
|
| 31 |
+
ALTER ROLE nlp_writer SET search_path = public;
|
| 32 |
+
|
| 33 |
+
\connect nlp_vectors
|
| 34 |
+
|
| 35 |
+
CREATE EXTENSION IF NOT EXISTS vector;
|
| 36 |
+
|
| 37 |
+
REVOKE CREATE ON SCHEMA public FROM PUBLIC;
|
| 38 |
+
GRANT USAGE ON SCHEMA public TO nlp_owner, nlp_reader, nlp_writer;
|
| 39 |
+
|
| 40 |
+
CREATE TABLE IF NOT EXISTS public.place_embeddings (
|
| 41 |
+
external_id TEXT PRIMARY KEY,
|
| 42 |
+
document TEXT NOT NULL,
|
| 43 |
+
metadata JSONB NOT NULL DEFAULT '{}'::jsonb,
|
| 44 |
+
embedding VECTOR(16) NOT NULL,
|
| 45 |
+
content_hash TEXT NOT NULL,
|
| 46 |
+
embedding_model TEXT NOT NULL,
|
| 47 |
+
embedding_version TEXT NOT NULL,
|
| 48 |
+
is_active BOOLEAN NOT NULL DEFAULT true,
|
| 49 |
+
updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
|
| 50 |
+
);
|
| 51 |
+
|
| 52 |
+
CREATE TABLE IF NOT EXISTS public.post_embeddings (
|
| 53 |
+
external_id TEXT PRIMARY KEY,
|
| 54 |
+
document TEXT NOT NULL,
|
| 55 |
+
metadata JSONB NOT NULL DEFAULT '{}'::jsonb,
|
| 56 |
+
embedding VECTOR(16) NOT NULL,
|
| 57 |
+
content_hash TEXT NOT NULL,
|
| 58 |
+
embedding_model TEXT NOT NULL,
|
| 59 |
+
embedding_version TEXT NOT NULL,
|
| 60 |
+
is_active BOOLEAN NOT NULL DEFAULT true,
|
| 61 |
+
updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
|
| 62 |
+
);
|
| 63 |
+
|
| 64 |
+
CREATE INDEX IF NOT EXISTS place_embeddings_embedding_hnsw_idx
|
| 65 |
+
ON public.place_embeddings USING hnsw (embedding vector_cosine_ops);
|
| 66 |
+
|
| 67 |
+
CREATE INDEX IF NOT EXISTS post_embeddings_embedding_hnsw_idx
|
| 68 |
+
ON public.post_embeddings USING hnsw (embedding vector_cosine_ops);
|
| 69 |
+
|
| 70 |
+
CREATE INDEX IF NOT EXISTS place_embeddings_metadata_gin_idx
|
| 71 |
+
ON public.place_embeddings USING gin (metadata);
|
| 72 |
+
|
| 73 |
+
CREATE INDEX IF NOT EXISTS post_embeddings_metadata_gin_idx
|
| 74 |
+
ON public.post_embeddings USING gin (metadata);
|
| 75 |
+
|
| 76 |
+
CREATE OR REPLACE FUNCTION public.match_places(
|
| 77 |
+
query_embedding VECTOR(16),
|
| 78 |
+
match_count INTEGER,
|
| 79 |
+
filters JSONB DEFAULT '{}'::jsonb
|
| 80 |
+
)
|
| 81 |
+
RETURNS TABLE (
|
| 82 |
+
external_id TEXT,
|
| 83 |
+
document TEXT,
|
| 84 |
+
metadata JSONB,
|
| 85 |
+
score DOUBLE PRECISION
|
| 86 |
+
)
|
| 87 |
+
LANGUAGE sql
|
| 88 |
+
STABLE
|
| 89 |
+
SECURITY DEFINER
|
| 90 |
+
SET search_path = public
|
| 91 |
+
AS $$
|
| 92 |
+
SELECT
|
| 93 |
+
p.external_id,
|
| 94 |
+
p.document,
|
| 95 |
+
p.metadata,
|
| 96 |
+
1 - (p.embedding <=> query_embedding) AS score
|
| 97 |
+
FROM public.place_embeddings p
|
| 98 |
+
WHERE p.is_active = true
|
| 99 |
+
AND COALESCE((filters->>'is_active')::boolean, true) = true
|
| 100 |
+
AND ((filters ? 'city') IS FALSE OR lower(p.metadata->>'city') = lower(filters->>'city'))
|
| 101 |
+
AND ((filters ? 'state') IS FALSE OR lower(p.metadata->>'state') = lower(filters->>'state'))
|
| 102 |
+
AND ((filters ? 'category') IS FALSE OR lower(p.metadata->>'category') = lower(filters->>'category'))
|
| 103 |
+
AND ((filters ? 'price_range') IS FALSE OR p.metadata->>'price_range' = filters->>'price_range')
|
| 104 |
+
AND ((filters ? 'occasion') IS FALSE OR p.metadata->>'occasion' ILIKE ('%' || (filters->>'occasion') || '%'))
|
| 105 |
+
ORDER BY p.embedding <=> query_embedding
|
| 106 |
+
LIMIT match_count;
|
| 107 |
+
$$;
|
| 108 |
+
|
| 109 |
+
CREATE OR REPLACE FUNCTION public.match_posts(
|
| 110 |
+
query_embedding VECTOR(16),
|
| 111 |
+
match_count INTEGER,
|
| 112 |
+
filters JSONB DEFAULT '{}'::jsonb
|
| 113 |
+
)
|
| 114 |
+
RETURNS TABLE (
|
| 115 |
+
external_id TEXT,
|
| 116 |
+
document TEXT,
|
| 117 |
+
metadata JSONB,
|
| 118 |
+
score DOUBLE PRECISION
|
| 119 |
+
)
|
| 120 |
+
LANGUAGE sql
|
| 121 |
+
STABLE
|
| 122 |
+
SECURITY DEFINER
|
| 123 |
+
SET search_path = public
|
| 124 |
+
AS $$
|
| 125 |
+
SELECT
|
| 126 |
+
p.external_id,
|
| 127 |
+
p.document,
|
| 128 |
+
p.metadata,
|
| 129 |
+
1 - (p.embedding <=> query_embedding) AS score
|
| 130 |
+
FROM public.post_embeddings p
|
| 131 |
+
WHERE p.is_active = true
|
| 132 |
+
AND COALESCE((filters->>'is_active')::boolean, true) = true
|
| 133 |
+
AND ((filters ? 'city') IS FALSE OR lower(p.metadata->>'city') = lower(filters->>'city'))
|
| 134 |
+
ORDER BY p.embedding <=> query_embedding
|
| 135 |
+
LIMIT match_count;
|
| 136 |
+
$$;
|
| 137 |
+
|
| 138 |
+
CREATE OR REPLACE FUNCTION public.upsert_place_embedding(
|
| 139 |
+
p_external_id TEXT,
|
| 140 |
+
p_document TEXT,
|
| 141 |
+
p_metadata JSONB,
|
| 142 |
+
p_embedding VECTOR(16),
|
| 143 |
+
p_content_hash TEXT,
|
| 144 |
+
p_embedding_model TEXT,
|
| 145 |
+
p_embedding_version TEXT,
|
| 146 |
+
p_is_active BOOLEAN
|
| 147 |
+
)
|
| 148 |
+
RETURNS VOID
|
| 149 |
+
LANGUAGE sql
|
| 150 |
+
SECURITY DEFINER
|
| 151 |
+
SET search_path = public
|
| 152 |
+
AS $$
|
| 153 |
+
INSERT INTO public.place_embeddings (
|
| 154 |
+
external_id,
|
| 155 |
+
document,
|
| 156 |
+
metadata,
|
| 157 |
+
embedding,
|
| 158 |
+
content_hash,
|
| 159 |
+
embedding_model,
|
| 160 |
+
embedding_version,
|
| 161 |
+
is_active,
|
| 162 |
+
updated_at
|
| 163 |
+
)
|
| 164 |
+
VALUES (
|
| 165 |
+
p_external_id,
|
| 166 |
+
p_document,
|
| 167 |
+
p_metadata,
|
| 168 |
+
p_embedding,
|
| 169 |
+
p_content_hash,
|
| 170 |
+
p_embedding_model,
|
| 171 |
+
p_embedding_version,
|
| 172 |
+
p_is_active,
|
| 173 |
+
now()
|
| 174 |
+
)
|
| 175 |
+
ON CONFLICT (external_id) DO UPDATE SET
|
| 176 |
+
document = EXCLUDED.document,
|
| 177 |
+
metadata = EXCLUDED.metadata,
|
| 178 |
+
embedding = EXCLUDED.embedding,
|
| 179 |
+
content_hash = EXCLUDED.content_hash,
|
| 180 |
+
embedding_model = EXCLUDED.embedding_model,
|
| 181 |
+
embedding_version = EXCLUDED.embedding_version,
|
| 182 |
+
is_active = EXCLUDED.is_active,
|
| 183 |
+
updated_at = now();
|
| 184 |
+
$$;
|
| 185 |
+
|
| 186 |
+
CREATE OR REPLACE FUNCTION public.upsert_post_embedding(
|
| 187 |
+
p_external_id TEXT,
|
| 188 |
+
p_document TEXT,
|
| 189 |
+
p_metadata JSONB,
|
| 190 |
+
p_embedding VECTOR(16),
|
| 191 |
+
p_content_hash TEXT,
|
| 192 |
+
p_embedding_model TEXT,
|
| 193 |
+
p_embedding_version TEXT,
|
| 194 |
+
p_is_active BOOLEAN
|
| 195 |
+
)
|
| 196 |
+
RETURNS VOID
|
| 197 |
+
LANGUAGE sql
|
| 198 |
+
SECURITY DEFINER
|
| 199 |
+
SET search_path = public
|
| 200 |
+
AS $$
|
| 201 |
+
INSERT INTO public.post_embeddings (
|
| 202 |
+
external_id,
|
| 203 |
+
document,
|
| 204 |
+
metadata,
|
| 205 |
+
embedding,
|
| 206 |
+
content_hash,
|
| 207 |
+
embedding_model,
|
| 208 |
+
embedding_version,
|
| 209 |
+
is_active,
|
| 210 |
+
updated_at
|
| 211 |
+
)
|
| 212 |
+
VALUES (
|
| 213 |
+
p_external_id,
|
| 214 |
+
p_document,
|
| 215 |
+
p_metadata,
|
| 216 |
+
p_embedding,
|
| 217 |
+
p_content_hash,
|
| 218 |
+
p_embedding_model,
|
| 219 |
+
p_embedding_version,
|
| 220 |
+
p_is_active,
|
| 221 |
+
now()
|
| 222 |
+
)
|
| 223 |
+
ON CONFLICT (external_id) DO UPDATE SET
|
| 224 |
+
document = EXCLUDED.document,
|
| 225 |
+
metadata = EXCLUDED.metadata,
|
| 226 |
+
embedding = EXCLUDED.embedding,
|
| 227 |
+
content_hash = EXCLUDED.content_hash,
|
| 228 |
+
embedding_model = EXCLUDED.embedding_model,
|
| 229 |
+
embedding_version = EXCLUDED.embedding_version,
|
| 230 |
+
is_active = EXCLUDED.is_active,
|
| 231 |
+
updated_at = now();
|
| 232 |
+
$$;
|
| 233 |
+
|
| 234 |
+
CREATE OR REPLACE FUNCTION public.get_place_content_hashes(p_external_ids TEXT[])
|
| 235 |
+
RETURNS TABLE (
|
| 236 |
+
external_id TEXT,
|
| 237 |
+
content_hash TEXT
|
| 238 |
+
)
|
| 239 |
+
LANGUAGE sql
|
| 240 |
+
STABLE
|
| 241 |
+
SECURITY DEFINER
|
| 242 |
+
SET search_path = public
|
| 243 |
+
AS $$
|
| 244 |
+
SELECT p.external_id, p.content_hash
|
| 245 |
+
FROM public.place_embeddings p
|
| 246 |
+
WHERE p.external_id = ANY(p_external_ids);
|
| 247 |
+
$$;
|
| 248 |
+
|
| 249 |
+
CREATE OR REPLACE FUNCTION public.get_post_content_hashes(p_external_ids TEXT[])
|
| 250 |
+
RETURNS TABLE (
|
| 251 |
+
external_id TEXT,
|
| 252 |
+
content_hash TEXT
|
| 253 |
+
)
|
| 254 |
+
LANGUAGE sql
|
| 255 |
+
STABLE
|
| 256 |
+
SECURITY DEFINER
|
| 257 |
+
SET search_path = public
|
| 258 |
+
AS $$
|
| 259 |
+
SELECT p.external_id, p.content_hash
|
| 260 |
+
FROM public.post_embeddings p
|
| 261 |
+
WHERE p.external_id = ANY(p_external_ids);
|
| 262 |
+
$$;
|
| 263 |
+
|
| 264 |
+
ALTER TABLE public.place_embeddings OWNER TO nlp_owner;
|
| 265 |
+
ALTER TABLE public.post_embeddings OWNER TO nlp_owner;
|
| 266 |
+
|
| 267 |
+
ALTER FUNCTION public.match_places(vector, integer, jsonb) OWNER TO nlp_owner;
|
| 268 |
+
ALTER FUNCTION public.match_posts(vector, integer, jsonb) OWNER TO nlp_owner;
|
| 269 |
+
ALTER FUNCTION public.upsert_place_embedding(text, text, jsonb, vector, text, text, text, boolean) OWNER TO nlp_owner;
|
| 270 |
+
ALTER FUNCTION public.upsert_post_embedding(text, text, jsonb, vector, text, text, text, boolean) OWNER TO nlp_owner;
|
| 271 |
+
ALTER FUNCTION public.get_place_content_hashes(text[]) OWNER TO nlp_owner;
|
| 272 |
+
ALTER FUNCTION public.get_post_content_hashes(text[]) OWNER TO nlp_owner;
|
| 273 |
+
|
| 274 |
+
REVOKE ALL ON public.place_embeddings FROM PUBLIC;
|
| 275 |
+
REVOKE ALL ON public.post_embeddings FROM PUBLIC;
|
| 276 |
+
REVOKE ALL ON FUNCTION public.match_places(vector, integer, jsonb) FROM PUBLIC;
|
| 277 |
+
REVOKE ALL ON FUNCTION public.match_posts(vector, integer, jsonb) FROM PUBLIC;
|
| 278 |
+
REVOKE ALL ON FUNCTION public.upsert_place_embedding(text, text, jsonb, vector, text, text, text, boolean) FROM PUBLIC;
|
| 279 |
+
REVOKE ALL ON FUNCTION public.upsert_post_embedding(text, text, jsonb, vector, text, text, text, boolean) FROM PUBLIC;
|
| 280 |
+
REVOKE ALL ON FUNCTION public.get_place_content_hashes(text[]) FROM PUBLIC;
|
| 281 |
+
REVOKE ALL ON FUNCTION public.get_post_content_hashes(text[]) FROM PUBLIC;
|
| 282 |
+
|
| 283 |
+
GRANT EXECUTE ON FUNCTION public.match_places(vector, integer, jsonb) TO nlp_reader;
|
| 284 |
+
GRANT EXECUTE ON FUNCTION public.match_posts(vector, integer, jsonb) TO nlp_reader;
|
| 285 |
+
|
| 286 |
+
GRANT EXECUTE ON FUNCTION public.upsert_place_embedding(text, text, jsonb, vector, text, text, text, boolean) TO nlp_writer;
|
| 287 |
+
GRANT EXECUTE ON FUNCTION public.upsert_post_embedding(text, text, jsonb, vector, text, text, text, boolean) TO nlp_writer;
|
| 288 |
+
GRANT EXECUTE ON FUNCTION public.get_place_content_hashes(text[]) TO nlp_writer;
|
| 289 |
+
GRANT EXECUTE ON FUNCTION public.get_post_content_hashes(text[]) TO nlp_writer;
|
| 290 |
+
|
| 291 |
+
SELECT to_regprocedure('public.match_places(vector, integer, jsonb)') AS match_places_signature;
|
| 292 |
+
SELECT to_regprocedure('public.match_posts(vector, integer, jsonb)') AS match_posts_signature;
|
| 293 |
+
SELECT has_function_privilege('nlp_reader', 'public.match_places(vector, integer, jsonb)', 'EXECUTE') AS reader_can_match_places;
|
| 294 |
+
SELECT has_function_privilege('nlp_reader', 'public.match_posts(vector, integer, jsonb)', 'EXECUTE') AS reader_can_match_posts;
|
tests/conftest.py
CHANGED
|
@@ -1,3 +1,4 @@
|
|
| 1 |
import os
|
| 2 |
|
| 3 |
os.environ["VECTOR_STORE_PROVIDER"] = "mock"
|
|
|
|
|
|
| 1 |
import os
|
| 2 |
|
| 3 |
os.environ["VECTOR_STORE_PROVIDER"] = "mock"
|
| 4 |
+
os.environ["GROQ_API_KEY"] = ""
|
tests/test_api_endpoints.py
CHANGED
|
@@ -44,6 +44,27 @@ def test_places_chat_endpoint_returns_trace_and_structured_places() -> None:
|
|
| 44 |
assert payload["metadata"]["places_used_as_context"]
|
| 45 |
|
| 46 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 47 |
def test_posts_recommendations_endpoint() -> None:
|
| 48 |
client = TestClient(create_app())
|
| 49 |
|
|
|
|
| 44 |
assert payload["metadata"]["places_used_as_context"]
|
| 45 |
|
| 46 |
|
| 47 |
+
def test_places_recommendations_returns_llm_message_and_tfidf_metadata() -> None:
|
| 48 |
+
client = TestClient(create_app())
|
| 49 |
+
|
| 50 |
+
response = client.post(
|
| 51 |
+
"/places/recommendations",
|
| 52 |
+
json={
|
| 53 |
+
"query": "quiero ver el atardecer y tomar fotos",
|
| 54 |
+
"city": "Tuxtla Gutierrez",
|
| 55 |
+
"filters": {"is_active": True},
|
| 56 |
+
"limit": 3,
|
| 57 |
+
},
|
| 58 |
+
)
|
| 59 |
+
|
| 60 |
+
assert response.status_code == 200
|
| 61 |
+
payload = response.json()
|
| 62 |
+
assert payload["message"]
|
| 63 |
+
assert payload["places"]
|
| 64 |
+
assert payload["metadata"]["ranking"] == "tfidf_cosine"
|
| 65 |
+
assert payload["metadata"]["used_llm"] is True
|
| 66 |
+
|
| 67 |
+
|
| 68 |
def test_posts_recommendations_endpoint() -> None:
|
| 69 |
client = TestClient(create_app())
|
| 70 |
|
tests/test_places_use_cases.py
CHANGED
|
@@ -3,10 +3,11 @@ from typing import Any, Sequence
|
|
| 3 |
import pytest
|
| 4 |
|
| 5 |
from app.modules.places.application.use_cases.chat_places import ChatPlacesUseCase
|
|
|
|
| 6 |
from app.modules.places.application.use_cases.search_places import SearchPlacesUseCase
|
| 7 |
from app.modules.places.domain.models import PlaceFilters
|
| 8 |
from app.modules.places.infrastructure.mock_place_repository import MockPlaceVectorRepository
|
| 9 |
-
from app.modules.places.infrastructure.
|
| 10 |
from app.shared.nlp.embeddings.mock import MockEmbeddingProvider
|
| 11 |
from app.shared.nlp.llm.base import LLMProvider, LLMResult
|
| 12 |
from app.shared.nlp.llm.mock import MockLLMProvider
|
|
@@ -19,7 +20,7 @@ async def test_search_places_use_case_with_mock_providers() -> None:
|
|
| 19 |
use_case = SearchPlacesUseCase(
|
| 20 |
embedding_provider=embedding_provider,
|
| 21 |
place_repository=MockPlaceVectorRepository(embedding_provider),
|
| 22 |
-
ranker=
|
| 23 |
)
|
| 24 |
|
| 25 |
result = await use_case.execute(
|
|
@@ -32,13 +33,32 @@ async def test_search_places_use_case_with_mock_providers() -> None:
|
|
| 32 |
assert all(place.city == "Tuxtla Gutierrez" for place in result.places)
|
| 33 |
|
| 34 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 35 |
@pytest.mark.asyncio
|
| 36 |
async def test_chat_places_returns_structured_places_with_llm_mock() -> None:
|
| 37 |
embedding_provider = MockEmbeddingProvider()
|
| 38 |
search_use_case = SearchPlacesUseCase(
|
| 39 |
embedding_provider=embedding_provider,
|
| 40 |
place_repository=MockPlaceVectorRepository(embedding_provider),
|
| 41 |
-
ranker=
|
| 42 |
)
|
| 43 |
chat_use_case = ChatPlacesUseCase(
|
| 44 |
search_use_case=search_use_case,
|
|
@@ -59,13 +79,41 @@ async def test_chat_places_returns_structured_places_with_llm_mock() -> None:
|
|
| 59 |
assert result.metadata["used_llm"] is True
|
| 60 |
|
| 61 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 62 |
@pytest.mark.asyncio
|
| 63 |
async def test_chat_places_falls_back_when_llm_fails() -> None:
|
| 64 |
embedding_provider = MockEmbeddingProvider()
|
| 65 |
search_use_case = SearchPlacesUseCase(
|
| 66 |
embedding_provider=embedding_provider,
|
| 67 |
place_repository=MockPlaceVectorRepository(embedding_provider),
|
| 68 |
-
ranker=
|
| 69 |
)
|
| 70 |
chat_use_case = ChatPlacesUseCase(
|
| 71 |
search_use_case=search_use_case,
|
|
@@ -95,3 +143,25 @@ class FailingLLMProvider(LLMProvider):
|
|
| 95 |
places: Sequence[dict[str, Any]],
|
| 96 |
) -> LLMResult:
|
| 97 |
raise RuntimeError("LLM failed")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 3 |
import pytest
|
| 4 |
|
| 5 |
from app.modules.places.application.use_cases.chat_places import ChatPlacesUseCase
|
| 6 |
+
from app.modules.places.application.use_cases.recommend_places import RecommendPlacesUseCase
|
| 7 |
from app.modules.places.application.use_cases.search_places import SearchPlacesUseCase
|
| 8 |
from app.modules.places.domain.models import PlaceFilters
|
| 9 |
from app.modules.places.infrastructure.mock_place_repository import MockPlaceVectorRepository
|
| 10 |
+
from app.modules.places.infrastructure.tfidf_place_ranker import TfidfPlaceRanker
|
| 11 |
from app.shared.nlp.embeddings.mock import MockEmbeddingProvider
|
| 12 |
from app.shared.nlp.llm.base import LLMProvider, LLMResult
|
| 13 |
from app.shared.nlp.llm.mock import MockLLMProvider
|
|
|
|
| 20 |
use_case = SearchPlacesUseCase(
|
| 21 |
embedding_provider=embedding_provider,
|
| 22 |
place_repository=MockPlaceVectorRepository(embedding_provider),
|
| 23 |
+
ranker=TfidfPlaceRanker(),
|
| 24 |
)
|
| 25 |
|
| 26 |
result = await use_case.execute(
|
|
|
|
| 33 |
assert all(place.city == "Tuxtla Gutierrez" for place in result.places)
|
| 34 |
|
| 35 |
|
| 36 |
+
@pytest.mark.asyncio
|
| 37 |
+
async def test_search_places_ranks_with_tfidf_cosine_similarity() -> None:
|
| 38 |
+
embedding_provider = MockEmbeddingProvider()
|
| 39 |
+
use_case = SearchPlacesUseCase(
|
| 40 |
+
embedding_provider=embedding_provider,
|
| 41 |
+
place_repository=MockPlaceVectorRepository(embedding_provider),
|
| 42 |
+
ranker=TfidfPlaceRanker(),
|
| 43 |
+
)
|
| 44 |
+
|
| 45 |
+
result = await use_case.execute(
|
| 46 |
+
query="atardecer fotos paseo",
|
| 47 |
+
filters=PlaceFilters(is_active=True),
|
| 48 |
+
limit=3,
|
| 49 |
+
)
|
| 50 |
+
|
| 51 |
+
assert result.places[0].id == "place_2"
|
| 52 |
+
assert result.places[0].score > 0
|
| 53 |
+
|
| 54 |
+
|
| 55 |
@pytest.mark.asyncio
|
| 56 |
async def test_chat_places_returns_structured_places_with_llm_mock() -> None:
|
| 57 |
embedding_provider = MockEmbeddingProvider()
|
| 58 |
search_use_case = SearchPlacesUseCase(
|
| 59 |
embedding_provider=embedding_provider,
|
| 60 |
place_repository=MockPlaceVectorRepository(embedding_provider),
|
| 61 |
+
ranker=TfidfPlaceRanker(),
|
| 62 |
)
|
| 63 |
chat_use_case = ChatPlacesUseCase(
|
| 64 |
search_use_case=search_use_case,
|
|
|
|
| 79 |
assert result.metadata["used_llm"] is True
|
| 80 |
|
| 81 |
|
| 82 |
+
@pytest.mark.asyncio
|
| 83 |
+
async def test_recommend_places_calls_llm_and_returns_message() -> None:
|
| 84 |
+
embedding_provider = MockEmbeddingProvider()
|
| 85 |
+
search_use_case = SearchPlacesUseCase(
|
| 86 |
+
embedding_provider=embedding_provider,
|
| 87 |
+
place_repository=MockPlaceVectorRepository(embedding_provider),
|
| 88 |
+
ranker=TfidfPlaceRanker(),
|
| 89 |
+
)
|
| 90 |
+
llm_provider = SpyLLMProvider()
|
| 91 |
+
use_case = RecommendPlacesUseCase(
|
| 92 |
+
search_use_case=search_use_case,
|
| 93 |
+
llm_provider=llm_provider,
|
| 94 |
+
output_guard=PlaceChatOutputGuard(),
|
| 95 |
+
)
|
| 96 |
+
|
| 97 |
+
result = await use_case.execute(
|
| 98 |
+
query="quiero una cena tranquila",
|
| 99 |
+
filters=PlaceFilters(city="Tuxtla Gutierrez", is_active=True),
|
| 100 |
+
limit=3,
|
| 101 |
+
)
|
| 102 |
+
|
| 103 |
+
assert llm_provider.calls == 1
|
| 104 |
+
assert result.message
|
| 105 |
+
assert result.places
|
| 106 |
+
assert result.metadata["used_llm"] is True
|
| 107 |
+
assert result.metadata["ranking"] == "tfidf_cosine"
|
| 108 |
+
|
| 109 |
+
|
| 110 |
@pytest.mark.asyncio
|
| 111 |
async def test_chat_places_falls_back_when_llm_fails() -> None:
|
| 112 |
embedding_provider = MockEmbeddingProvider()
|
| 113 |
search_use_case = SearchPlacesUseCase(
|
| 114 |
embedding_provider=embedding_provider,
|
| 115 |
place_repository=MockPlaceVectorRepository(embedding_provider),
|
| 116 |
+
ranker=TfidfPlaceRanker(),
|
| 117 |
)
|
| 118 |
chat_use_case = ChatPlacesUseCase(
|
| 119 |
search_use_case=search_use_case,
|
|
|
|
| 143 |
places: Sequence[dict[str, Any]],
|
| 144 |
) -> LLMResult:
|
| 145 |
raise RuntimeError("LLM failed")
|
| 146 |
+
|
| 147 |
+
|
| 148 |
+
class SpyLLMProvider(LLMProvider):
|
| 149 |
+
provider_name = "spy-llama"
|
| 150 |
+
model_name = "spy-model"
|
| 151 |
+
|
| 152 |
+
def __init__(self) -> None:
|
| 153 |
+
self.calls = 0
|
| 154 |
+
|
| 155 |
+
async def generate_place_chat_response(
|
| 156 |
+
self,
|
| 157 |
+
user_intent: str,
|
| 158 |
+
region: str | None,
|
| 159 |
+
places: Sequence[dict[str, Any]],
|
| 160 |
+
) -> LLMResult:
|
| 161 |
+
self.calls += 1
|
| 162 |
+
place_names = ", ".join(str(place["name"]) for place in places[:2])
|
| 163 |
+
return LLMResult(
|
| 164 |
+
message=f"Estas opciones pueden encajar con tu plan: {place_names}.",
|
| 165 |
+
provider=self.provider_name,
|
| 166 |
+
model=self.model_name,
|
| 167 |
+
)
|