feat:根据登陆人角色反馈回答

This commit is contained in:
shijing 2026-07-15 13:18:37 +08:00
parent b128b3c2ac
commit e89bc896a6
3 changed files with 416 additions and 29 deletions

View File

@ -24,6 +24,7 @@ from app.schemas.chat import (
MessageResponse, MessageResponse,
) )
from app.services.audit_service import write_audit from app.services.audit_service import write_audit
from app.services.document_service import resolve_chat_doc_scope
def _attachment_data_url(att: Attachment) -> str | None: def _attachment_data_url(att: Attachment) -> str | None:
@ -172,8 +173,33 @@ def stream_chat(
db: Session = Depends(get_db), db: Session = Depends(get_db),
) -> StreamingResponse: ) -> StreamingResponse:
current_user_id = current_user.id current_user_id = current_user.id
is_admin = current_user.is_superuser or any(
role.code in {"admin", "system_admin"} for role in current_user.roles
)
allowed_tags = [] if current_user.is_superuser else [role.code for role in current_user.roles] allowed_tags = [] if current_user.is_superuser else [role.code for role in current_user.roles]
if current_user.is_superuser:
role_labels: list[str] = ["系统管理员"]
else:
role_labels = [role.name for role in current_user.roles if role.name]
role_hint = "".join(role_labels)
# 按用户所属单位(含逐级向上回退)计算允许引用的文档 title 集合。
# 超管/管理员或用户没绑定单位 -> 不做过滤allowed_titles=None
scope = {"allowed_titles": None, "own_unit": "", "chosen_unit": "", "fell_back": False, "public_only": False}
if not is_admin and current_user.unit is not None:
try:
adp_docs = (AdpClient().list_documents(page=1, page_size=200).get("list") or [])
scope = resolve_chat_doc_scope(current_user, adp_docs)
except Exception:
# ADP 拉取失败时不阻断问答,退化为不过滤
scope = {"allowed_titles": None, "own_unit": (current_user.unit.name or "").strip(),
"chosen_unit": "", "fell_back": False, "public_only": False}
elif current_user.unit is not None:
scope["own_unit"] = (current_user.unit.name or "").strip()
allowed_titles = scope["allowed_titles"]
if not (payload.question or "").strip() and not payload.attachment_ids: if not (payload.question or "").strip() and not payload.attachment_ids:
raise HTTPException(status_code=400, detail="问题文本和图片至少需要一项") raise HTTPException(status_code=400, detail="问题文本和图片至少需要一项")
@ -220,45 +246,94 @@ def stream_chat(
att.message_id = message_id att.message_id = message_id
db.commit() db.commit()
raw_question = payload.question or ""
hint_lines: list[str] = []
if role_hint:
hint_lines.append(
f"【提问人职位:{role_hint}"
f"请优先按与该职位相关的制度、职责与权限范围给出答案;"
f"若该职位不涉及相关业务,请明确说明并给出通用回答。"
)
own_unit = scope.get("own_unit") or ""
chosen_unit = scope.get("chosen_unit") or ""
if own_unit and not is_admin:
if scope.get("fell_back") and chosen_unit:
hint_lines.append(
f"【提问人所属单位:{own_unit}"
f"该单位无直接对应制度,请引用其上级单位「{chosen_unit}」的对应制度以及公开制度作答。"
)
elif chosen_unit:
hint_lines.append(
f"【提问人所属单位:{own_unit}"
f"请优先引用本单位对应制度及公开制度作答,不要引用其他单位专属制度。"
)
elif scope.get("public_only"):
hint_lines.append(
f"【提问人所属单位:{own_unit}"
f"当前单位链上均无对应制度,仅可引用公开制度作答,若信息不足请明确说明。"
)
adp_question = ("\n".join(hint_lines) + "\n\n" + raw_question) if hint_lines else raw_question
def event_stream(): def event_stream():
answer_parts: list[str] = [] answer_parts: list[str] = []
usage = {"token_usage": 0, "pu_usage": 0} usage = {"token_usage": 0, "pu_usage": 0}
references: list[dict] = [] references: list[dict] = []
seen_refs: set[tuple[str, str]] = set() seen_refs: set[tuple[str, str]] = set()
for event in AdpClient().stream_answer( stream_start = datetime.utcnow()
payload.question, failed = False
user_id=str(current_user_id), error_message: str | None = None
conversation_id=adp_conversation_id, try:
allowed_tags=allowed_tags, for event in AdpClient().stream_answer(
image_urls=image_urls, adp_question,
): user_id=str(current_user_id),
event_name = event["event"] conversation_id=adp_conversation_id,
if event_name == "answer.delta": allowed_tags=allowed_tags,
answer_parts.append(event["data"]) image_urls=image_urls,
elif event_name == "answer.reference": ):
title = (event["data"].get("source_title") or "").strip() event_name = event["event"]
if not title: if event_name == "answer.delta":
continue answer_parts.append(event["data"])
key = (title, (event["data"].get("snippet") or "").strip()) elif event_name == "answer.reference":
if key in seen_refs: title = (event["data"].get("source_title") or "").strip()
continue if not title:
seen_refs.add(key) continue
references.append(event["data"]) # 按用户所属单位链过滤引用allowed_titles=None 时不过滤
elif event_name == "answer.usage": if allowed_titles is not None and title not in allowed_titles:
usage.update(event["data"]) continue
yield f"event: {event_name}\ndata: {json.dumps(event['data'], ensure_ascii=False)}\n\n" key = (title, (event["data"].get("snippet") or "").strip())
if key in seen_refs:
continue
seen_refs.add(key)
references.append(event["data"])
elif event_name == "answer.usage":
usage.update(event["data"])
yield f"event: {event_name}\ndata: {json.dumps(event['data'], ensure_ascii=False)}\n\n"
except Exception as exc: # noqa: BLE001
failed = True
error_message = str(exc) or exc.__class__.__name__
yield f"event: answer.error\ndata: {json.dumps({'message': error_message}, ensure_ascii=False)}\n\n"
latency_ms = int((datetime.utcnow() - stream_start).total_seconds() * 1000)
save_db = SessionLocal() save_db = SessionLocal()
try: try:
saved = save_db.get(Message, message_id) saved = save_db.get(Message, message_id)
if saved: if saved:
saved.answer = "".join(answer_parts) saved.answer = "".join(answer_parts)
saved.status = "done" saved.status = "failed" if failed else "done"
saved.token_usage = usage["token_usage"] saved.token_usage = usage["token_usage"]
saved.pu_usage = usage["pu_usage"] saved.pu_usage = usage["pu_usage"]
for ref in references: saved.latency_ms = latency_ms
save_db.add(MessageReference(message_id=message_id, **ref)) if not failed:
write_audit(save_db, "chat.ask", actor_id=current_user_id, target_type="message", target_id=str(message_id)) for ref in references:
save_db.add(MessageReference(message_id=message_id, **ref))
write_audit(
save_db,
"chat.ask.failed" if failed else "chat.ask",
actor_id=current_user_id,
target_type="message",
target_id=str(message_id),
detail={"error": error_message} if failed and error_message else None,
)
save_db.commit() save_db.commit()
finally: finally:
save_db.close() save_db.close()

View File

@ -1,15 +1,70 @@
from fastapi import APIRouter, Depends from datetime import datetime, timedelta
from fastapi import APIRouter, Depends, Query
from sqlalchemy import func, select from sqlalchemy import func, select
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from app.db import get_db from app.db import get_db
from app.deps import require_admin from app.deps import require_admin
from app.models import AuditLog, Document, Message, User from app.models import (
ApiCallLog,
AuditLog,
Document,
Message,
MessageReference,
User,
)
from app.schemas.audit import StatsOverviewResponse from app.schemas.audit import StatsOverviewResponse
router = APIRouter() router = APIRouter()
_RANGE_SPEC: dict[str, tuple[timedelta, int, str]] = {
"1h": (timedelta(hours=1), 12, "5min"),
"12h": (timedelta(hours=12), 12, "hour"),
"24h": (timedelta(hours=24), 24, "hour"),
"7d": (timedelta(days=7), 7, "day"),
}
def _resolve_range(range_code: str, now: datetime) -> tuple[datetime, datetime, int, str]:
span, bucket_count, unit = _RANGE_SPEC[range_code]
return now - span, now, bucket_count, unit
def _bucket_index(dt: datetime, start: datetime, unit: str) -> int | None:
if dt < start:
return None
delta = (dt - start).total_seconds()
if unit == "5min":
return int(delta // 300)
if unit == "hour":
return int(delta // 3600)
if unit == "day":
return int(delta // 86400)
return None
def _bucket_label(start: datetime, idx: int, unit: str) -> str:
if unit == "5min":
return (start + timedelta(minutes=5 * idx)).strftime("%H:%M")
if unit == "hour":
return (start + timedelta(hours=idx)).strftime("%H:00")
if unit == "day":
return (start + timedelta(days=idx)).strftime("%m-%d")
return str(idx)
def _pct(part: int, total: int) -> float:
return round(part / total * 100, 1) if total else 0.0
def _delta_pct(current: int, previous: int) -> float | None:
if previous == 0:
return None
return round((current - previous) / previous * 100, 1)
@router.get("/overview", response_model=StatsOverviewResponse) @router.get("/overview", response_model=StatsOverviewResponse)
def overview(db: Session = Depends(get_db), _: User = Depends(require_admin)) -> StatsOverviewResponse: def overview(db: Session = Depends(get_db), _: User = Depends(require_admin)) -> StatsOverviewResponse:
question_count = db.scalar(select(func.count(Message.id))) or 0 question_count = db.scalar(select(func.count(Message.id))) or 0
@ -26,3 +81,188 @@ def overview(db: Session = Depends(get_db), _: User = Depends(require_admin)) ->
token_usage=token_usage, token_usage=token_usage,
pu_usage=pu_usage, pu_usage=pu_usage,
) )
@router.get("/dashboard/business")
def dashboard_business(
range_code: str = Query("24h", alias="range", pattern="^(1h|12h|24h|7d)$"),
db: Session = Depends(get_db),
_: User = Depends(require_admin),
) -> dict:
now = datetime.utcnow()
start, end, bucket_count, unit = _resolve_range(range_code, now)
prev_start = start - (end - start)
rows = db.execute(
select(
Message.user_id,
Message.status,
Message.latency_ms,
Message.created_at,
).where(Message.created_at >= start, Message.created_at < end)
).all()
prev_user_count = (
db.scalar(
select(func.count(func.distinct(Message.user_id))).where(
Message.created_at >= prev_start, Message.created_at < start
)
)
or 0
)
total = len(rows)
success = sum(1 for r in rows if r.status == "done")
failed = sum(1 for r in rows if r.status == "failed")
latencies = [r.latency_ms for r in rows if r.latency_ms]
active_users = {r.user_id for r in rows}
active_user_count = len(active_users)
trend = [
{"bucket": _bucket_label(start, i, unit), "asked": 0, "success": 0, "failed": 0}
for i in range(bucket_count)
]
for r in rows:
idx = _bucket_index(r.created_at, start, unit)
if idx is None or idx >= bucket_count:
continue
trend[idx]["asked"] += 1
if r.status == "done":
trend[idx]["success"] += 1
elif r.status == "failed":
trend[idx]["failed"] += 1
top_rows = db.execute(
select(Message.question, func.count(Message.id).label("cnt"))
.where(Message.created_at >= start, Message.created_at < end)
.group_by(Message.question)
.order_by(func.count(Message.id).desc())
.limit(10)
).all()
top_questions = [{"question": (r[0] or "").strip()[:120], "count": int(r[1])} for r in top_rows]
return {
"range": {
"code": range_code,
"start": start.isoformat(),
"end": end.isoformat(),
"bucket_unit": unit,
},
"kpis": {
"question_count": total,
"active_user_count": active_user_count,
"prev_active_user_count": int(prev_user_count),
"user_delta_pct": _delta_pct(active_user_count, int(prev_user_count)),
"avg_questions_per_user": round(total / active_user_count, 1) if active_user_count else 0,
"avg_latency_ms": int(sum(latencies) / len(latencies)) if latencies else 0,
"success_count": success,
"success_ratio": _pct(success, total),
"fail_count": failed,
"fail_ratio": _pct(failed, total),
},
"trend": trend,
"top_questions": top_questions,
}
@router.get("/dashboard/resources")
def dashboard_resources(
range_code: str = Query("24h", alias="range", pattern="^(1h|12h|24h|7d)$"),
db: Session = Depends(get_db),
_: User = Depends(require_admin),
) -> dict:
now = datetime.utcnow()
start, end, bucket_count, unit = _resolve_range(range_code, now)
document_count = (
db.scalar(select(func.count(Document.id)).where(Document.status != "deleted")) or 0
)
total_tokens = (
db.scalar(
select(func.coalesce(func.sum(Message.token_usage), 0)).where(
Message.created_at >= start, Message.created_at < end
)
)
or 0
)
total_pu = (
db.scalar(
select(func.coalesce(func.sum(Message.pu_usage), 0)).where(
Message.created_at >= start, Message.created_at < end
)
)
or 0
)
api_total = (
db.scalar(
select(func.count(ApiCallLog.id)).where(
ApiCallLog.created_at >= start, ApiCallLog.created_at < end
)
)
or 0
)
api_success = (
db.scalar(
select(func.count(ApiCallLog.id)).where(
ApiCallLog.created_at >= start,
ApiCallLog.created_at < end,
ApiCallLog.status == "success",
)
)
or 0
)
api_avg_latency = db.scalar(
select(func.avg(ApiCallLog.latency_ms)).where(
ApiCallLog.created_at >= start,
ApiCallLog.created_at < end,
ApiCallLog.latency_ms.is_not(None),
)
)
msg_rows = db.execute(
select(Message.created_at, Message.token_usage, Message.pu_usage).where(
Message.created_at >= start, Message.created_at < end
)
).all()
token_trend = [
{"bucket": _bucket_label(start, i, unit), "tokens": 0, "pu": 0}
for i in range(bucket_count)
]
for r in msg_rows:
idx = _bucket_index(r.created_at, start, unit)
if idx is None or idx >= bucket_count:
continue
token_trend[idx]["tokens"] += int(r.token_usage or 0)
token_trend[idx]["pu"] += int(r.pu_usage or 0)
ref_rows = db.execute(
select(MessageReference.source_title, func.count(MessageReference.id).label("cnt"))
.join(Message, Message.id == MessageReference.message_id)
.where(Message.created_at >= start, Message.created_at < end)
.group_by(MessageReference.source_title)
.order_by(func.count(MessageReference.id).desc())
.limit(10)
).all()
top_documents = [
{"title": (r[0] or "").strip()[:120], "count": int(r[1])} for r in ref_rows
]
return {
"range": {
"code": range_code,
"start": start.isoformat(),
"end": end.isoformat(),
"bucket_unit": unit,
},
"kpis": {
"document_count": int(document_count),
"total_tokens": int(total_tokens),
"total_pu": int(total_pu),
"api_call_count": int(api_total),
"api_success_ratio": _pct(int(api_success), int(api_total)),
"avg_api_latency_ms": int(api_avg_latency) if api_avg_latency else 0,
},
"token_trend": token_trend,
"top_documents": top_documents,
}

View File

@ -1,7 +1,79 @@
from sqlalchemy import select from sqlalchemy import select
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from app.models import Document, DocumentTag from app.models import Document, DocumentTag, User
def _extract_doc_labels(doc: dict) -> set[str]:
labels: set[str] = set()
for attr in (doc.get("attr_labels") or []):
for lbl in (attr.get("labels") or []):
name = (lbl.get("label_name") or "").strip()
if name:
labels.add(name.replace("可见", ""))
return labels
def build_unit_chain(user: User) -> list[str]:
"""从用户自身单位一路向上到根节点,返回单位名列表。"""
chain: list[str] = []
seen: set[int] = set()
cur = user.unit
while cur is not None and cur.id not in seen:
seen.add(cur.id)
name = (cur.name or "").strip()
if name:
chain.append(name)
cur = cur.parent
return chain
def resolve_chat_doc_scope(user: User, adp_docs: list[dict]) -> dict:
""""就近单位优先、逐级向上回退"的规则计算问答阶段允许引用的文档集合。
- 超管/管理员不做过滤
- 普通用户从自身单位开始找到最近一层"有直接匹配文档"的单位
allowed = 该层匹配文档 所有"所有人可见"文档
若整条链都没有匹配单位的文档 allowed = 仅公开文档
"""
is_admin = user.is_superuser or any(
r.code in {"admin", "system_admin"} for r in user.roles
)
unit_chain = build_unit_chain(user)
result = {
"allowed_titles": None, # None 表示不过滤
"unit_chain": unit_chain,
"own_unit": unit_chain[0] if unit_chain else "",
"chosen_unit": "", # 实际命中的单位名,可能是祖先单位
"fell_back": False, # True = 走到了祖先单位
"public_only": False, # True = 整条链无匹配,仅公开
}
if is_admin or not unit_chain:
return result
public: set[str] = set()
per_unit: dict[str, set[str]] = {u: set() for u in unit_chain}
for doc in adp_docs:
title = (doc.get("file_name") or "").strip()
if not title:
continue
bases = _extract_doc_labels(doc)
if "所有人" in bases:
public.add(title)
for u in unit_chain:
if any(u in b for b in bases):
per_unit[u].add(title)
for idx, u in enumerate(unit_chain):
if per_unit[u]:
result["allowed_titles"] = public | per_unit[u]
result["chosen_unit"] = u
result["fell_back"] = idx > 0
return result
result["allowed_titles"] = public
result["public_only"] = True
return result
def get_or_create_tags(db: Session, names: list[str]) -> list[DocumentTag]: def get_or_create_tags(db: Session, names: list[str]) -> list[DocumentTag]: