feat(avatar): add user token accounting
This commit is contained in:
@@ -16,12 +16,20 @@ import embeddings
|
||||
from database import get_db
|
||||
from models import Avatar, KnowledgeChunk, KnowledgeDoc, QAPair, User
|
||||
from responses import ok, fail
|
||||
from services.token_billing import (
|
||||
InsufficientTokensError,
|
||||
estimate_fallback_usage,
|
||||
release_reservation,
|
||||
reserve_avatar_tokens,
|
||||
settle_reservation,
|
||||
)
|
||||
|
||||
router = APIRouter(tags=["数字分身聊天"])
|
||||
|
||||
CHAT_API_URL = os.getenv("CHAT_API_URL", "https://dashscope.aliyuncs.com/compatible-mode/v1")
|
||||
CHAT_API_KEY = os.getenv("CHAT_API_KEY", "")
|
||||
CHAT_MODEL = os.getenv("CHAT_MODEL", "qwen-plus")
|
||||
CHAT_MAX_OUTPUT_TOKENS = max(128, int(os.getenv("CHAT_MAX_OUTPUT_TOKENS", "1024")))
|
||||
MAX_MESSAGE_LENGTH = 4000
|
||||
MAX_HISTORY_MESSAGES = 10
|
||||
QA_LEXICAL_THRESHOLD = 0.72
|
||||
@@ -279,7 +287,7 @@ def _search_knowledge(db: Session, avatar_id: str, question: str, top_k: int = 5
|
||||
return results
|
||||
|
||||
|
||||
def _call_qwen(messages: list[dict], temperature: float) -> str:
|
||||
def _call_qwen(messages: list[dict], temperature: float) -> dict:
|
||||
if not CHAT_API_KEY:
|
||||
raise RuntimeError("Qwen 模型服务未配置 CHAT_API_KEY")
|
||||
url = f"{CHAT_API_URL.rstrip('/')}/chat/completions"
|
||||
@@ -287,6 +295,7 @@ def _call_qwen(messages: list[dict], temperature: float) -> str:
|
||||
"model": CHAT_MODEL,
|
||||
"messages": messages,
|
||||
"temperature": temperature,
|
||||
"max_tokens": CHAT_MAX_OUTPUT_TOKENS,
|
||||
}
|
||||
try:
|
||||
response = httpx.post(
|
||||
@@ -302,7 +311,7 @@ def _call_qwen(messages: list[dict], temperature: float) -> str:
|
||||
raise RuntimeError("Qwen 模型服务暂时不可用") from exc
|
||||
if not isinstance(answer, str) or not answer.strip():
|
||||
raise RuntimeError("Qwen 模型没有返回有效回答")
|
||||
return answer.strip()
|
||||
return {"answer": answer.strip(), "usage": data.get("usage") or {}}
|
||||
|
||||
|
||||
def _iter_qwen_stream(messages: list[dict], temperature: float):
|
||||
@@ -310,7 +319,14 @@ def _iter_qwen_stream(messages: list[dict], temperature: float):
|
||||
if not CHAT_API_KEY:
|
||||
raise RuntimeError("模型服务未配置")
|
||||
url = f"{CHAT_API_URL.rstrip('/')}/chat/completions"
|
||||
payload = {"model": CHAT_MODEL, "messages": messages, "temperature": temperature, "stream": True}
|
||||
payload = {
|
||||
"model": CHAT_MODEL,
|
||||
"messages": messages,
|
||||
"temperature": temperature,
|
||||
"max_tokens": CHAT_MAX_OUTPUT_TOKENS,
|
||||
"stream": True,
|
||||
"stream_options": {"include_usage": True},
|
||||
}
|
||||
try:
|
||||
with httpx.stream("POST", url, headers={"Authorization": f"Bearer {CHAT_API_KEY}"}, json=payload, timeout=45) as response:
|
||||
response.raise_for_status()
|
||||
@@ -322,11 +338,15 @@ def _iter_qwen_stream(messages: list[dict], temperature: float):
|
||||
if data == "[DONE]":
|
||||
return
|
||||
try:
|
||||
delta = json.loads(data).get("choices", [{}])[0].get("delta", {}).get("content")
|
||||
parsed = json.loads(data)
|
||||
except (ValueError, IndexError, AttributeError):
|
||||
continue
|
||||
if parsed.get("usage"):
|
||||
yield {"usage": parsed["usage"]}
|
||||
choices = parsed.get("choices") or []
|
||||
delta = choices[0].get("delta", {}).get("content") if choices else None
|
||||
if delta:
|
||||
yield delta
|
||||
yield {"content": delta}
|
||||
except httpx.HTTPError as exc:
|
||||
raise RuntimeError("模型服务暂时不可用") from exc
|
||||
|
||||
@@ -350,6 +370,7 @@ def _resolve_reply(
|
||||
qa_pairs: list[Any] | None = None,
|
||||
search_fn: Callable[..., list[dict]] | None = None,
|
||||
model_client: Callable[..., str] | None = None,
|
||||
usage_source: str = "chat",
|
||||
) -> dict:
|
||||
if qa_pairs is None:
|
||||
qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all()
|
||||
@@ -362,16 +383,49 @@ def _resolve_reply(
|
||||
messages = _build_prompt(avatar, history, question, hits)
|
||||
config = _config(avatar)
|
||||
temperature = min(0.45 if hits else 0.25, 0.2 + config["creativity"] / 100 * 0.6)
|
||||
model_client = model_client or _call_qwen
|
||||
answer = model_client(messages=messages, temperature=temperature)
|
||||
return {
|
||||
token_usage = None
|
||||
if model_client is not None:
|
||||
answer = model_client(messages=messages, temperature=temperature)
|
||||
else:
|
||||
reservation = reserve_avatar_tokens(
|
||||
db,
|
||||
avatar,
|
||||
usage_source,
|
||||
CHAT_MODEL,
|
||||
messages,
|
||||
CHAT_MAX_OUTPUT_TOKENS,
|
||||
)
|
||||
try:
|
||||
model_result = _call_qwen(messages=messages, temperature=temperature)
|
||||
answer = model_result["answer"]
|
||||
token_usage = settle_reservation(
|
||||
db,
|
||||
reservation,
|
||||
model_result.get("usage"),
|
||||
fallback_total=estimate_fallback_usage(messages, answer),
|
||||
)
|
||||
except Exception as exc:
|
||||
release_reservation(db, reservation, str(exc))
|
||||
raise
|
||||
result = {
|
||||
"answer": answer,
|
||||
"source": "knowledge" if hits else "qwen",
|
||||
"references": hits,
|
||||
}
|
||||
if token_usage:
|
||||
result["tokenUsage"] = token_usage
|
||||
return result
|
||||
|
||||
|
||||
def _stream_reply(db: Session, avatar: Avatar, question: str, history: list[Any], *, public: bool = False):
|
||||
def _stream_reply(
|
||||
db: Session,
|
||||
avatar: Avatar,
|
||||
question: str,
|
||||
history: list[Any],
|
||||
*,
|
||||
public: bool = False,
|
||||
usage_source: str = "chat_stream",
|
||||
):
|
||||
qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all()
|
||||
matched = _match_standard_qa(question, qa_pairs)
|
||||
if matched:
|
||||
@@ -381,18 +435,62 @@ def _stream_reply(db: Session, avatar: Avatar, question: str, history: list[Any]
|
||||
source = "knowledge" if references else "qwen"
|
||||
config = _config(avatar)
|
||||
temperature = min(0.45 if references else 0.25, 0.2 + config["creativity"] / 100 * 0.6)
|
||||
chunks = _iter_qwen_stream(_build_prompt(avatar, history, question, references), temperature)
|
||||
messages = _build_prompt(avatar, history, question, references)
|
||||
reservation = reserve_avatar_tokens(
|
||||
db,
|
||||
avatar,
|
||||
usage_source,
|
||||
CHAT_MODEL,
|
||||
messages,
|
||||
CHAT_MAX_OUTPUT_TOKENS,
|
||||
)
|
||||
chunks = _iter_qwen_stream(messages, temperature)
|
||||
if matched:
|
||||
messages, reservation = [], None
|
||||
if public:
|
||||
source, references = "public", []
|
||||
|
||||
def generate():
|
||||
output_parts = []
|
||||
provider_usage = None
|
||||
settled = False
|
||||
try:
|
||||
yield _sse("meta", {"source": source, "references": references})
|
||||
for content in chunks:
|
||||
for chunk in chunks:
|
||||
if reservation is None:
|
||||
content = chunk
|
||||
else:
|
||||
provider_usage = chunk.get("usage") or provider_usage
|
||||
content = chunk.get("content")
|
||||
if not content:
|
||||
continue
|
||||
output_parts.append(content)
|
||||
yield _sse("delta", {"content": content})
|
||||
yield _sse("done", {})
|
||||
token_usage = None
|
||||
if reservation is not None:
|
||||
answer = "".join(output_parts)
|
||||
token_usage = settle_reservation(
|
||||
db,
|
||||
reservation,
|
||||
provider_usage,
|
||||
fallback_total=estimate_fallback_usage(messages, answer),
|
||||
)
|
||||
settled = True
|
||||
yield _sse("done", {} if public else {"tokenUsage": token_usage})
|
||||
except RuntimeError as exc:
|
||||
yield _sse("error", {"message": str(exc)})
|
||||
finally:
|
||||
if reservation is not None and not settled:
|
||||
answer = "".join(output_parts)
|
||||
if answer:
|
||||
settle_reservation(
|
||||
db,
|
||||
reservation,
|
||||
provider_usage,
|
||||
fallback_total=estimate_fallback_usage(messages, answer),
|
||||
)
|
||||
else:
|
||||
release_reservation(db, reservation, "stream_ended_without_output")
|
||||
|
||||
return StreamingResponse(
|
||||
generate(),
|
||||
@@ -441,11 +539,14 @@ def get_shared_avatar(share_token: str, db: Session = Depends(get_db)):
|
||||
def public_chat(share_token: str, body: ChatIn = Body(...), db: Session = Depends(get_db)):
|
||||
avatar = _require_shared_avatar(db, share_token)
|
||||
try:
|
||||
result = _resolve_reply(db, avatar, body.message, body.history)
|
||||
result = _resolve_reply(db, avatar, body.message, body.history, usage_source="public_chat")
|
||||
# 公开访客无需获知知识文件名、检索分数或内部答复来源。
|
||||
result["references"] = []
|
||||
result["source"] = "public"
|
||||
result.pop("tokenUsage", None)
|
||||
return ok(result)
|
||||
except InsufficientTokensError as exc:
|
||||
return fail(str(exc), code=402)
|
||||
except RuntimeError as exc:
|
||||
return fail(str(exc), code=502)
|
||||
|
||||
@@ -455,15 +556,30 @@ def chat(avatar_id: str, body: ChatIn = Body(...), authorization: str = Header(N
|
||||
avatar = _require_owned_avatar(db, avatar_id, authorization)
|
||||
try:
|
||||
return ok(_resolve_reply(db, avatar, body.message, body.history))
|
||||
except InsufficientTokensError as exc:
|
||||
return fail(str(exc), code=402)
|
||||
except RuntimeError as exc:
|
||||
return fail(str(exc), code=502)
|
||||
|
||||
|
||||
@router.post("/avatar/{avatar_id}/chat/stream")
|
||||
def chat_stream(avatar_id: str, body: ChatIn = Body(...), authorization: str = Header(None), db: Session = Depends(get_db)):
|
||||
return _stream_reply(db, _require_owned_avatar(db, avatar_id, authorization), body.message, body.history)
|
||||
try:
|
||||
return _stream_reply(db, _require_owned_avatar(db, avatar_id, authorization), body.message, body.history)
|
||||
except InsufficientTokensError as exc:
|
||||
raise HTTPException(status_code=402, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.post("/public/avatar/{share_token}/chat/stream")
|
||||
def public_chat_stream(share_token: str, body: ChatIn = Body(...), db: Session = Depends(get_db)):
|
||||
return _stream_reply(db, _require_shared_avatar(db, share_token), body.message, body.history, public=True)
|
||||
try:
|
||||
return _stream_reply(
|
||||
db,
|
||||
_require_shared_avatar(db, share_token),
|
||||
body.message,
|
||||
body.history,
|
||||
public=True,
|
||||
usage_source="public_chat_stream",
|
||||
)
|
||||
except InsufficientTokensError as exc:
|
||||
raise HTTPException(status_code=402, detail=str(exc)) from exc
|
||||
|
||||
Reference in New Issue
Block a user