Files
xiaobaifupan/next/backend/features/review/service.py
T

369 lines
16 KiB
Python

from __future__ import annotations
import secrets
from collections.abc import Iterator
from dataclasses import dataclass
from datetime import date, datetime
from typing import Any
from zoneinfo import ZoneInfo
from backend.data.gateway import DataGateway
from backend.database.connection import Database
from backend.features.accounts.models import Principal
from backend.features.review.prompt import PROMPT_VERSION, messages
from backend.features.review.repository import ReviewRepository
from backend.features.review.views import alert as alert_view
from backend.features.review.views import note as note_view
from backend.features.review.views import trade_summary, trades
from backend.features.screener.service import ScreenerService
from backend.llm.gateway import LLMCall, LLMGateway, LLMGatewayError
SHANGHAI = ZoneInfo("Asia/Shanghai")
class ReviewError(RuntimeError):
pass
@dataclass(frozen=True, slots=True)
class PreparedAssistant:
call: LLMCall
prompt: list[dict[str, str]]
context_date: str
class ReviewService:
def __init__(
self,
database: Database,
repository: ReviewRepository,
gateway: DataGateway,
screener: ScreenerService,
llm: LLMGateway,
) -> None:
self._database = database
self._repository = repository
self._gateway = gateway
self._screener = screener
self._llm = llm
def workspace(self, principal: Principal, requested_date: str) -> dict[str, Any]:
trade_date = self._trade_date(requested_date)
with self._database.read() as connection:
watches = self._watch_rows(connection, principal.user.id, trade_date)
daily = self._repository.note(connection, principal.user.id, "", trade_date)
history = self._repository.notes(connection, principal.user.id, "", 60)
trade_rows = trades(self._repository.trades(connection, principal.user.id))
return {
"trade_date": trade_date,
"watchlist": watches,
"daily": note_view(daily),
"history": [note_view(row) for row in history],
"trades": trade_rows,
"trade_summary": trade_summary(trade_rows),
}
def add_watch(self, principal: Principal, identifier: str) -> dict[str, Any]:
normalized = identifier.strip().upper()
candidates = self._gateway.search(normalized)
item = next(
(
candidate
for candidate in candidates
if candidate.entity_type == "stock" and candidate.identifier == normalized
),
None,
)
if item is None:
raise ReviewError("未找到可加入自选的股票。")
with self._database.transaction() as connection:
self._repository.save_watch(
connection,
principal.user.id,
{"identifier": item.identifier, "name": item.name, "sector": item.sector},
_now(),
)
return {"identifier": item.identifier, "name": item.name}
def save_watch_remark(self, principal: Principal, identifier: str, remark: str) -> None:
normalized = _text(remark, 500)
with self._database.transaction() as connection:
if not self._repository.save_watch_remark(
connection, principal.user.id, identifier, normalized
):
raise ReviewError("自选记录不存在。")
def delete_watch(self, principal: Principal, identifier: str) -> None:
with self._database.transaction() as connection:
if not self._repository.delete_watch(connection, principal.user.id, identifier):
raise ReviewError("自选记录不存在。")
def save_note(self, principal: Principal, payload: dict[str, Any]) -> int:
data = {
"trade_date": _date(payload["trade_date"]),
"code": payload.get("code", "").strip().upper(),
"stock_name": _text(payload.get("stock_name", ""), 40),
"summary": _text(payload.get("summary", ""), 500),
"content": _text(payload.get("content", ""), 5000),
"plan": _text(payload.get("plan", ""), 2000),
}
if data["code"] and not data["stock_name"]:
raise ReviewError("个股名称不能为空。")
with self._database.transaction() as connection:
return self._repository.save_note(connection, principal.user.id, data, _now())
def stock_notes(self, principal: Principal, code: str) -> list[dict[str, Any]]:
with self._database.read() as connection:
return [
note_view(row)
for row in self._repository.notes(
connection, principal.user.id, code.strip().upper()
)
]
def delete_note(self, principal: Principal, note_id: int) -> None:
with self._database.transaction() as connection:
if not self._repository.delete_note(connection, principal.user.id, note_id):
raise ReviewError("复盘记录不存在。")
def save_trade(self, principal: Principal, payload: dict[str, Any]) -> int:
data = dict(payload)
data["trade_date"] = _date(str(data["trade_date"]))
data["code"] = str(data["code"]).strip().upper()
data["name"] = _text(data["name"], 40, True)
data["thesis"] = _text(data.get("thesis", ""), 2000)
data["execution"] = _text(data.get("execution", ""), 2000)
data["tags"] = [_text(item, 20, True) for item in data.get("tags", [])][:8]
with self._database.transaction() as connection:
return self._repository.save_trade(connection, principal.user.id, data, _now())
def delete_trade(self, principal: Principal, trade_id: int) -> None:
with self._database.transaction() as connection:
if not self._repository.delete_trade(connection, principal.user.id, trade_id):
raise ReviewError("交易记录不存在。")
def alert_center(self, principal: Principal, unread: bool = False) -> dict[str, Any]:
self._sync_tracking_alerts(principal.user.id)
today = date.today().isoformat()
with self._database.read() as connection:
rows = self._repository.alerts(connection, principal.user.id, unread)
items = [alert_view(row, today) for row in rows]
if unread:
items = [item for item in items if item["due"]]
count = self._repository.unread_count(connection, principal.user.id, today)
return {"items": items, "unread_count": count, "as_of": today}
def create_alert(self, principal: Principal, payload: dict[str, Any]) -> int:
data = {
"kind": "manual",
"title": _text(payload["title"], 80, True),
"content": _text(payload.get("content", ""), 500),
"available_date": _date(payload["remind_date"]),
"code": _text(payload.get("code", ""), 12),
"dedupe_key": f"manual:{secrets.token_hex(12)}",
}
with self._database.transaction() as connection:
return self._repository.save_alert(connection, principal.user.id, data, _now())
def mark_alert(self, principal: Principal, alert_id: int) -> None:
with self._database.transaction() as connection:
if not self._repository.mark_alert(connection, principal.user.id, alert_id, _now()):
raise ReviewError("提醒不存在。")
def mark_all_alerts(self, principal: Principal) -> int:
with self._database.transaction() as connection:
return self._repository.mark_all_alerts(
connection, principal.user.id, date.today().isoformat(), _now()
)
def delete_alert(self, principal: Principal, alert_id: int) -> None:
with self._database.transaction() as connection:
if not self._repository.delete_alert(connection, principal.user.id, alert_id):
raise ReviewError("提醒不存在。")
def messages(self, principal: Principal) -> list[dict[str, Any]]:
with self._database.read() as connection:
return [dict(row) for row in self._repository.messages(connection, principal.user.id)]
def clear_messages(self, principal: Principal) -> int:
with self._database.transaction() as connection:
return self._repository.clear_messages(connection, principal.user.id)
def prepare_assistant(
self, principal: Principal, requested_date: str, question: str
) -> PreparedAssistant:
context_date = self._trade_date(requested_date)
normalized = _text(question, 2000, True)
context = self._assistant_context(principal, context_date)
history = self._history(principal.user.id)
prompt = messages(context, history, normalized)
call = self._llm.prepare(
principal,
feature="review_assistant",
prompt_version=PROMPT_VERSION,
business_id=context_date,
input_chars=sum(len(item["content"]) for item in prompt),
)
with self._database.transaction() as connection:
self._repository.add_message(
connection,
principal.user.id,
{
"role": "user",
"content": normalized,
"context_date": context_date,
"request_id": call.request_id,
"status": "complete",
},
_now(),
)
return PreparedAssistant(call, prompt, context_date)
def stream_assistant(self, prepared: PreparedAssistant) -> Iterator[dict[str, Any]]:
answer = ""
stream = self._llm.stream(prepared.call, prepared.prompt)
try:
for event in stream:
if event.type == "delta":
answer += event.content
yield {
"type": "delta",
"content": event.content,
"request_id": event.request_id,
}
elif event.type == "done":
self._save_assistant(prepared, answer, "complete")
yield {"type": "done", "request_id": event.request_id}
except GeneratorExit:
stream.close()
if answer:
self._save_assistant(prepared, answer, "stopped")
raise
except LLMGatewayError as exc:
if answer:
self._save_assistant(prepared, answer, "error")
yield {"type": "error", "code": exc.code, "message": str(exc), "partial": exc.partial}
def _watch_rows(self, connection, user_id: int, trade_date: str) -> list[dict[str, Any]]:
factors = self._repository.latest_factors(connection, trade_date)
auction_scores = self._repository.latest_auction_scores(connection, trade_date)
rows = []
for watch in self._repository.watchlist(connection, user_id):
factor = factors.get(str(watch["identifier"]), {})
rows.append(
{
**dict(watch),
"code": str(watch["identifier"]).split(".")[0],
"pct_chg": factor.get("pct_chg"),
"return_5d": factor.get("return_5d"),
"attention_score": auction_scores.get(str(watch["identifier"])),
"remark": str(watch["remark"]),
}
)
return rows
def _sync_tracking_alerts(self, user_id: int) -> None:
tracks = self._screener.tracks(user_id)
batches: dict[int, list[dict]] = {}
for item in tracks:
batches.setdefault(int(item["run_id"]), []).append(item)
with self._database.transaction() as connection:
for run_id, items in batches.items():
name = str(items[0]["strategy_name"])
if any(item.get("t1_return") is not None for item in items):
self._repository.save_alert(
connection,
user_id,
{
"kind": "strategy_t1",
"title": f"{name} 已有 T+1 反馈",
"content": (
f"{sum(item.get('t1_return') is not None for item in items)}"
f"/{len(items)} 只标的已有首日表现。"
),
"available_date": date.today().isoformat(),
"code": "",
"dedupe_key": f"strategy:{run_id}:t1",
},
_now(),
)
if items and all(item.get("t5_return") is not None for item in items):
self._repository.save_alert(
connection,
user_id,
{
"kind": "strategy_t5",
"title": f"{name} 五日跟踪完成",
"content": f"本批 {len(items)} 只标的已完成 T+5 跟踪。",
"available_date": date.today().isoformat(),
"code": "",
"dedupe_key": f"strategy:{run_id}:t5",
},
_now(),
)
def _assistant_context(self, principal: Principal, trade_date: str) -> dict[str, Any]:
summary = self._gateway.summary(trade_date)
workspace = self.workspace(principal, trade_date)
alerts = self.alert_center(principal)["items"][:20]
return {
"data_date": trade_date,
"market_facts": summary["values"],
"user_records": {
"watchlist": workspace["watchlist"][:30],
"daily_reviews": workspace["history"][:10],
"strategy_tracking": self._screener.tracks(principal.user.id)[:30],
"alerts": alerts,
"trade_summary": workspace["trade_summary"],
"trade_entries": workspace["trades"][:30],
},
}
def _history(self, user_id: int) -> list[dict[str, str]]:
with self._database.read() as connection:
rows = self._repository.messages(connection, user_id, 20)
return [
{"role": str(row["role"]), "content": str(row["content"])[:4000]} for row in rows[-12:]
]
def _save_assistant(self, prepared: PreparedAssistant, answer: str, status: str) -> None:
if not answer:
return
with self._database.transaction() as connection:
self._repository.add_message(
connection,
prepared.call.user_id,
{
"role": "assistant",
"content": answer,
"context_date": prepared.context_date,
"request_id": prepared.call.request_id,
"status": status,
},
_now(),
)
def _trade_date(self, value: str) -> str:
requested = _date(value)
return self._gateway.trade_context(requested).actual_date or requested
def _date(value: str) -> str:
try:
return date.fromisoformat(value).isoformat()
except ValueError as exc:
raise ReviewError("日期格式无效。") from exc
def _text(value: Any, limit: int, required: bool = False) -> str:
normalized = " ".join(str(value or "").split())
if required and not normalized:
raise ReviewError("必填内容不能为空。")
if len(normalized) > limit:
raise ReviewError(f"内容不能超过{limit}个字符。")
return normalized
def _now() -> str:
return datetime.now(SHANGHAI).isoformat(timespec="seconds")