Skip to content
Merged
Show file tree
Hide file tree
Changes from all 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
@@ -1,10 +1,14 @@
APP_ENV=local
INTERNAL_API_KEY=change-me
SPRING_API_BASE_URL=http://localhost:8080
SUPABASE_CONNECT_TIMEOUT_SECONDS=5
SUPABASE_DB_URL=postgresql://user:password@host:5432/postgres
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