Skip to content
Merged
Show file tree
Hide file tree
Changes from 2 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions backend-ai/.env.example
Original file line number Diff line number Diff line change
Expand Up @@ -2,9 +2,13 @@ APP_ENV=local
INTERNAL_API_KEY=change-me
SPRING_API_BASE_URL=http://localhost:8080
SUPABASE_DB_URL=postgresql://user:password@host:5432/postgres
SUPABASE_CONNECT_TIMEOUT_SECONDS=5
SUPABASE_STATEMENT_TIMEOUT_MS=5000
SUPABASE_SERVICE_ROLE_KEY=
GMS_API_KEY=
LLM_MODEL=
EMBEDDING_API_KEY=
EMBEDDING_BASE_URL=https://api.openai.com/v1
EMBEDDING_MODEL=
LANGSMITH_TRACING=false
LANGSMITH_API_KEY=
45 changes: 45 additions & 0 deletions backend-ai/app/clients/embedding_client.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,45 @@
from typing import Any

import httpx

from app.core.config import get_settings


class EmbeddingClient:
"""HTTP boundary for query embeddings."""

def __init__(
self,
api_key: str | None = None,
base_url: str | None = None,
model: str | None = None,
timeout_seconds: float = 10.0,
) -> None:
settings = get_settings()
self.api_key = api_key if api_key is not None else settings.embedding_api_key
self.base_url = (base_url or settings.embedding_base_url).rstrip("/")
self.model = model or settings.embedding_model
self.timeout_seconds = timeout_seconds

def embed_query(self, query: str) -> list[float]:
if not query.strip():
raise ValueError("query must not be blank.")
if not self.api_key or not self.model:
raise RuntimeError("Embedding client is not configured.")

response = httpx.post(
f"{self.base_url}/embeddings",
headers={"Authorization": f"Bearer {self.api_key}"},
json={"model": self.model, "input": query},
timeout=self.timeout_seconds,
)
response.raise_for_status()
payload: dict[str, Any] = response.json()
data = payload.get("data")
if not isinstance(data, list) or not data or not isinstance(data[0], dict):
raise RuntimeError("Embedding response did not include data.")

embedding = data[0].get("embedding")
if not isinstance(embedding, list):
raise RuntimeError("Embedding response did not include an embedding vector.")
return [float(value) for value in embedding]
79 changes: 57 additions & 22 deletions backend-ai/app/clients/supabase_client.py
Original file line number Diff line number Diff line change
@@ -1,30 +1,65 @@
from typing import Any

import psycopg
from psycopg.rows import dict_row

from app.core.config import get_settings


class SupabaseVectorClient:
"""Client boundary for Supabase PostgreSQL + pgvector legal search."""

def __init__(self) -> None:
self.database_url = get_settings().supabase_db_url

def similarity_search_legal_documents(self, query: str, top_k: int = 3) -> list[dict[str, Any]]:
return [
{
"lawName": "주택임대차보호법",
"articleNo": "제3조",
"title": "대항력",
"source": "legal-stub",
"content": "임차인은 주택의 인도와 주민등록을 마친 때에는 그 다음 날부터 제3자에 대하여 효력이 생깁니다.",
"score": 0.91,
},
{
"lawName": "주택임대차보호법",
"articleNo": "제3조의2",
"title": "보증금의 회수",
"source": "legal-stub",
"content": "확정일자를 갖춘 임차인은 경매 또는 공매 시 후순위권리자보다 우선하여 보증금을 변제받을 수 있습니다.",
"score": 0.86,
},
][:top_k]
def __init__(
self,
connect_timeout_seconds: int | None = None,
statement_timeout_ms: int | None = None,
) -> None:
settings = get_settings()
self.database_url = settings.supabase_db_url
self.connect_timeout_seconds = (
connect_timeout_seconds
if connect_timeout_seconds is not None
else settings.supabase_connect_timeout_seconds
)
self.statement_timeout_ms = (
statement_timeout_ms
if statement_timeout_ms is not None
else settings.supabase_statement_timeout_ms
)

def similarity_search_legal_documents(
self,
query_embedding: list[float],
top_k: int = 3,
) -> list[dict[str, Any]]:
vector_literal = to_pgvector_literal(query_embedding)
sql = """
select
law_name,
article_no,
article_title,
content,
1 - (embedding <=> %s::vector) as score
from public.legal_document_chunks
where embedding is not null
order by embedding <=> %s::vector
limit %s
"""
with psycopg.connect(
self.database_url,
row_factory=dict_row,
connect_timeout=self.connect_timeout_seconds,
) as conn:
with conn.cursor() as cursor:
cursor.execute(
"set local statement_timeout = %s",
(self.statement_timeout_ms,),
)
cursor.execute(sql, (vector_literal, vector_literal, top_k))
return list(cursor.fetchall())


def to_pgvector_literal(embedding: list[float]) -> str:
if not embedding:
raise ValueError("query_embedding must not be empty.")
return "[" + ",".join(f"{float(value):.10g}" for value in embedding) + "]"
13 changes: 13 additions & 0 deletions backend-ai/app/core/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,9 +15,22 @@ class Settings(BaseSettings):
default="postgresql://user:password@host:5432/postgres",
alias="SUPABASE_DB_URL",
)
supabase_connect_timeout_seconds: int = Field(
default=5,
alias="SUPABASE_CONNECT_TIMEOUT_SECONDS",
)
supabase_statement_timeout_ms: int = Field(
default=5000,
alias="SUPABASE_STATEMENT_TIMEOUT_MS",
)
supabase_service_role_key: str = Field(default="", alias="SUPABASE_SERVICE_ROLE_KEY")
gms_api_key: str = Field(default="", alias="GMS_API_KEY")
llm_model: str = Field(default="", alias="LLM_MODEL")
embedding_api_key: str = Field(default="", alias="EMBEDDING_API_KEY")
embedding_base_url: str = Field(
default="https://api.openai.com/v1",
alias="EMBEDDING_BASE_URL",
)
embedding_model: str = Field(default="", alias="EMBEDDING_MODEL")
langsmith_tracing: bool = Field(default=False, alias="LANGSMITH_TRACING")
langsmith_api_key: str = Field(default="", alias="LANGSMITH_API_KEY")
Expand Down
2 changes: 1 addition & 1 deletion backend-ai/app/graph/nodes/legal_rag.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,6 @@ def legal_rag(state: AgentState) -> AgentState:
"legal_cards": cards,
"tool_results": {
**state.get("tool_results", {}),
"legalRag": {"topK": len(cards), "source": "supabase-pgvector-stub"},
"legalRag": {"topK": len(cards), "source": "supabase-pgvector"},
},
}
50 changes: 46 additions & 4 deletions backend-ai/app/rag/retriever.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,53 @@
from typing import Any
from typing import Any, Protocol

from app.clients.embedding_client import EmbeddingClient
from app.clients.supabase_client import SupabaseVectorClient


class EmbeddingClientProtocol(Protocol):
def embed_query(self, query: str) -> list[float]:
...


class VectorClientProtocol(Protocol):
def similarity_search_legal_documents(
self,
query_embedding: list[float],
top_k: int,
) -> list[dict[str, Any]]:
...


class LegalRetriever:
def __init__(self) -> None:
self.vector_client = SupabaseVectorClient()
def __init__(
self,
embedding_client: EmbeddingClientProtocol | None = None,
vector_client: VectorClientProtocol | None = None,
max_top_k: int = 5,
) -> None:
self.embedding_client = embedding_client or EmbeddingClient()
self.vector_client = vector_client or SupabaseVectorClient()
self.max_top_k = max_top_k

def retrieve(self, query: str, top_k: int = 3) -> list[dict[str, Any]]:
return self.vector_client.similarity_search_legal_documents(query=query, top_k=top_k)
normalized_query = query.strip()
if not normalized_query:
return []

safe_top_k = min(max(top_k, 1), self.max_top_k)
query_embedding = self.embedding_client.embed_query(normalized_query)
rows = self.vector_client.similarity_search_legal_documents(
query_embedding=query_embedding,
top_k=safe_top_k,
)
return [normalize_legal_card(row) for row in rows]


def normalize_legal_card(row: dict[str, Any]) -> dict[str, Any]:
return {
"lawName": row.get("lawName") or row["law_name"],
"articleNo": row.get("articleNo") or row["article_no"],
"title": row.get("title") or row["article_title"],
"content": row["content"],
"score": float(row["score"]),
}
29 changes: 23 additions & 6 deletions backend-ai/tests/test_agent_chat.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
from fastapi.testclient import TestClient

from app.api.routes import get_agent_graph
from app.core.config import get_settings
from app.graph.nodes import legal_rag as legal_rag_module
from app.graph.nodes.classify_intent import classify_message
from app.graph.state import Intent
from app.main import app
Expand Down Expand Up @@ -46,14 +48,29 @@ def test_agent_chat_returns_intent_and_answer() -> None:
assert "properties" in body


def test_agent_chat_returns_legal_cards_for_legal_question() -> None:
def test_agent_chat_returns_legal_cards_for_legal_question(monkeypatch) -> None:
class FakeRetriever:
def retrieve(self, query: str, top_k: int = 3) -> list[dict]:
return [
{
"lawName": "주택임대차보호법",
"articleNo": "제3조의2",
"title": "보증금의 회수",
"content": "임차인은 보증금을 우선변제받을 권리가 있다.",
"score": 0.86,
}
]

monkeypatch.setattr(legal_rag_module, "LegalRetriever", FakeRetriever)
get_agent_graph.cache_clear()

response = client.post(
"/internal/agent/chat",
headers=internal_api_headers(),
json={
"userId": "user-1",
"sessionId": None,
"message": "확정일자는 언제 받아야 하나요?",
"message": "계약 전 보증금 반환 관련 법을 알려줘",
"context": {"selectedPropertyId": None, "recentMessages": []},
},
)
Expand All @@ -73,7 +90,7 @@ def test_agent_chat_returns_legal_cards_for_legal_question() -> None:

def test_classify_intent_examples() -> None:
assert classify_message("관악구 보증금 5천 이하 원룸 추천해줘") == Intent.PROPERTY_SEARCH
assert classify_message("전세사기 계약이면 어떻게 해야 해?") == Intent.LEGAL_CONSULT
assert classify_message("이 매물 시세가 비싼 편이야?") == Intent.PRICE_ANALYSIS
assert classify_message("주변 치안과 CCTV는 괜찮아?") == Intent.SAFETY_ANALYSIS
assert classify_message("HUG 보증보험 가입 가능해?") == Intent.HUG_CALC
assert classify_message("계약 전에 법을 확인하고 싶어") == Intent.LEGAL_CONSULT
assert classify_message("이 매물 가격이 비싼 편이야?") == Intent.PRICE_ANALYSIS
assert classify_message("주변 cctv는 괜찮아?") == Intent.SAFETY_ANALYSIS
assert classify_message("hug 보증보험 가능해?") == Intent.HUG_CALC
Loading