Files

101 lines
3.8 KiB
Python
Raw Permalink Normal View History

import json
from datetime import UTC, datetime
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.sql.elements import ColumnElement
from app.core.knowledge_contracts import KnowledgeHit, KnowledgeQuery
from app.model.knowledge import FinKnowledgeMeta
class KnowledgeMysqlAuthority:
def __init__(self, session: AsyncSession) -> None:
self._session = session
async def filter_published(self, hits: tuple[KnowledgeHit, ...]) -> list[KnowledgeHit]:
ids = tuple(int(hit.knowledge_id) for hit in hits if hit.knowledge_id.isdecimal())
if not ids:
return []
rows = await self._session.scalars(
select(FinKnowledgeMeta).where(
FinKnowledgeMeta.id.in_(ids), *self._published_filters()
)
)
approved = {str(row.id): row for row in rows}
result: list[KnowledgeHit] = []
for hit in hits:
row = approved.get(hit.knowledge_id)
if row is None:
continue
answer = self.extract_answer(row.content_text).strip()
if answer:
result.append(hit.model_copy(update={
"answer": answer, "version": row.version, "title": row.title,
}))
return result
async def search_keyword(
self, query: KnowledgeQuery, collections: tuple[str, ...], top_k: int
) -> list[KnowledgeHit]:
"""向量服务不可用时,在原授权集合内执行受限的只读关键词检索。"""
keyword = self._keyword(query.query)
if not collections or not keyword:
return []
# 显式转义 LIKE 通配符,避免用户输入扩大关键词降级的匹配范围。
escaped_keyword = keyword.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
rows = await self._session.scalars(
select(FinKnowledgeMeta)
.where(
FinKnowledgeMeta.milvus_collection.in_(collections),
*self._published_filters(),
FinKnowledgeMeta.content_text.like(f"%{escaped_keyword}%", escape="\\"),
)
.order_by(FinKnowledgeMeta.id.desc())
.limit(top_k)
)
result: list[KnowledgeHit] = []
for row in rows:
answer = self.extract_answer(row.content_text).strip()
if not answer:
continue
result.append(
KnowledgeHit(
knowledge_id=str(row.id),
collection=row.milvus_collection,
title=row.title,
snippet=answer[:300],
answer=answer,
version=row.version,
)
)
return result
@staticmethod
def _keyword(query: str) -> str:
"""压缩空白并限制关键词长度,避免降级查询承载无界输入。"""
return "".join(query.split())[:64]
@staticmethod
def _published_filters() -> tuple[ColumnElement[bool], ...]:
today = datetime.now(UTC).date()
return (
FinKnowledgeMeta.review_status == "published",
FinKnowledgeMeta.status == "active",
(FinKnowledgeMeta.effective_date.is_(None))
| (FinKnowledgeMeta.effective_date <= today),
(FinKnowledgeMeta.expire_date.is_(None))
| (FinKnowledgeMeta.expire_date > today),
)
@staticmethod
def extract_answer(content_text: str) -> str:
try:
payload = json.loads(content_text)
except json.JSONDecodeError:
return content_text
answer = payload.get("answer") if isinstance(payload, dict) else None
if isinstance(answer, str):
return answer
return content_text