feat(avatar): add private vision chat support
This commit is contained in:
@@ -1,21 +1,32 @@
|
||||
import difflib
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import secrets
|
||||
import string
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Any, Callable
|
||||
|
||||
import httpx
|
||||
from fastapi import APIRouter, Body, Depends, Header, HTTPException
|
||||
from fastapi import APIRouter, Body, Depends, File, Header, HTTPException, UploadFile
|
||||
from fastapi.responses import StreamingResponse
|
||||
from pydantic import BaseModel, Field
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
import embeddings
|
||||
from database import get_db
|
||||
from models import Avatar, KnowledgeChunk, KnowledgeDoc, QAPair, User
|
||||
from models import Avatar, ChatAttachment, KnowledgeChunk, KnowledgeDoc, QAPair, User
|
||||
from responses import ok, fail
|
||||
from services.vision_service import (
|
||||
GENERAL_VISION_PROMPT,
|
||||
MEDICAL_OCR_PROMPT,
|
||||
ImageValidationError,
|
||||
build_attachment_warning,
|
||||
call_vision_model,
|
||||
parse_vision_analysis,
|
||||
prepare_image,
|
||||
)
|
||||
from services.token_billing import (
|
||||
InsufficientTokensError,
|
||||
estimate_fallback_usage,
|
||||
@@ -24,8 +35,10 @@ from services.token_billing import (
|
||||
settle_reservation,
|
||||
)
|
||||
from services.chat_model_config import ChatModelConfig, get_chat_model_config
|
||||
from services.chat_attachment_service import purge_expired_chat_attachments
|
||||
|
||||
router = APIRouter(tags=["数字分身聊天"])
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
MAX_MESSAGE_LENGTH = 4000
|
||||
MAX_HISTORY_MESSAGES = 10
|
||||
@@ -49,14 +62,35 @@ _KOREAN_HANGUL = re.compile(r"[\uac00-\ud7af\u1100-\u11ff]")
|
||||
|
||||
|
||||
class ChatMessage(BaseModel):
|
||||
model_config = ConfigDict(populate_by_name=True)
|
||||
|
||||
role: str = Field(pattern="^(user|assistant)$")
|
||||
content: str = Field(min_length=1, max_length=MAX_MESSAGE_LENGTH)
|
||||
attachment_ids: list[str] = Field(
|
||||
default_factory=list,
|
||||
alias="attachmentIds",
|
||||
max_length=3,
|
||||
)
|
||||
|
||||
|
||||
class ChatIn(BaseModel):
|
||||
message: str = Field(min_length=1, max_length=MAX_MESSAGE_LENGTH)
|
||||
model_config = ConfigDict(populate_by_name=True)
|
||||
|
||||
message: str = Field(default="", max_length=MAX_MESSAGE_LENGTH)
|
||||
attachment_ids: list[str] = Field(
|
||||
default_factory=list,
|
||||
alias="attachmentIds",
|
||||
max_length=3,
|
||||
)
|
||||
history: list[ChatMessage] = Field(default_factory=list, max_length=MAX_HISTORY_MESSAGES)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def require_message_or_image(self):
|
||||
self.message = self.message.strip()
|
||||
if not self.message and not self.attachment_ids:
|
||||
raise ValueError("请输入消息或选择图片")
|
||||
return self
|
||||
|
||||
|
||||
def _resolve_user(authorization: str | None, db: Session):
|
||||
if not authorization:
|
||||
@@ -77,6 +111,228 @@ def _require_owned_avatar(db: Session, avatar_id: str, authorization: str | None
|
||||
return avatar
|
||||
|
||||
|
||||
def _attachment_expiry() -> datetime:
|
||||
retention_hours = max(
|
||||
1, min(168, int(os.getenv("CHAT_ATTACHMENT_RETENTION_HOURS", "24")))
|
||||
)
|
||||
return datetime.utcnow() + timedelta(hours=retention_hours)
|
||||
|
||||
|
||||
def _chat_attachment_ids(body: ChatIn) -> list[str]:
|
||||
values = list(body.attachment_ids)
|
||||
for message in body.history[-MAX_HISTORY_MESSAGES:]:
|
||||
values.extend(message.attachment_ids)
|
||||
unique = list(dict.fromkeys(str(value).strip() for value in values if str(value).strip()))
|
||||
if len(unique) > 3:
|
||||
raise HTTPException(status_code=400, detail="一次会话最多引用 3 张图片")
|
||||
return unique
|
||||
|
||||
|
||||
def _load_chat_attachments(db: Session, avatar_id: str, body: ChatIn) -> list[ChatAttachment]:
|
||||
attachment_ids = _chat_attachment_ids(body)
|
||||
if not attachment_ids:
|
||||
return []
|
||||
purge_expired_chat_attachments(db)
|
||||
rows = db.query(ChatAttachment).filter(
|
||||
ChatAttachment.avatar_id == avatar_id,
|
||||
ChatAttachment.id.in_(attachment_ids),
|
||||
).all()
|
||||
by_id = {row.id: row for row in rows}
|
||||
if len(by_id) != len(attachment_ids):
|
||||
raise HTTPException(status_code=400, detail="图片资料不存在、已过期或不属于当前分身")
|
||||
ordered = [by_id[attachment_id] for attachment_id in attachment_ids]
|
||||
if any(row.status != "ready" for row in ordered):
|
||||
raise HTTPException(status_code=409, detail="图片尚未识别完成,请稍后重试")
|
||||
now = datetime.utcnow()
|
||||
for row in ordered:
|
||||
row.used_at = now
|
||||
db.commit()
|
||||
return ordered
|
||||
|
||||
|
||||
def _attachment_contexts(rows: list[ChatAttachment]) -> list[dict]:
|
||||
contexts = []
|
||||
remaining_text = 12000
|
||||
for row in rows:
|
||||
extracted = (row.extracted_text or "")[:remaining_text]
|
||||
remaining_text = max(0, remaining_text - len(extracted))
|
||||
contexts.append({
|
||||
"id": row.id,
|
||||
"filename": row.filename,
|
||||
"category": row.category,
|
||||
"summary": row.summary,
|
||||
"extractedText": extracted,
|
||||
"structuredData": row.structured_data or {},
|
||||
"warning": row.warning,
|
||||
})
|
||||
return contexts
|
||||
|
||||
|
||||
def _image_retrieval_question(question: str, image_contexts: list[dict]) -> str:
|
||||
parts = [question.strip()]
|
||||
for context in image_contexts:
|
||||
parts.extend([
|
||||
str(context.get("summary") or "")[:600],
|
||||
str(context.get("extractedText") or "")[:1200],
|
||||
])
|
||||
return "\n".join(part for part in parts if part).strip()
|
||||
|
||||
|
||||
def _run_billed_vision_call(
|
||||
db: Session,
|
||||
avatar: Avatar,
|
||||
prepared,
|
||||
*,
|
||||
model: str,
|
||||
prompt: str,
|
||||
source: str,
|
||||
json_output: bool,
|
||||
model_config: ChatModelConfig,
|
||||
) -> dict:
|
||||
estimate_messages = [{
|
||||
"role": "user",
|
||||
"content": f"[一张待识别图片]\n{prompt}",
|
||||
}]
|
||||
reservation = reserve_avatar_tokens(
|
||||
db,
|
||||
avatar,
|
||||
source,
|
||||
model,
|
||||
estimate_messages,
|
||||
model_config.vision_max_tokens,
|
||||
minimum_reserve_tokens=max(
|
||||
1000, int(os.getenv("VISION_TOKEN_RESERVE", "12000"))
|
||||
),
|
||||
)
|
||||
try:
|
||||
result = call_vision_model(
|
||||
prepared,
|
||||
model_config,
|
||||
model=model,
|
||||
prompt=prompt,
|
||||
json_output=json_output,
|
||||
)
|
||||
settle_reservation(
|
||||
db,
|
||||
reservation,
|
||||
result.get("usage"),
|
||||
fallback_total=estimate_fallback_usage(
|
||||
estimate_messages, result.get("content") or ""
|
||||
),
|
||||
)
|
||||
return result
|
||||
except Exception as exc:
|
||||
release_reservation(db, reservation, str(exc))
|
||||
raise
|
||||
|
||||
|
||||
async def _analyze_uploaded_image(
|
||||
db: Session,
|
||||
avatar: Avatar,
|
||||
file: UploadFile,
|
||||
*,
|
||||
uploader_kind: str,
|
||||
) -> ChatAttachment:
|
||||
max_bytes = max(1024, int(os.getenv("CHAT_IMAGE_MAX_BYTES", str(8 * 1024 * 1024))))
|
||||
content = await file.read(max_bytes + 1)
|
||||
filename = os.path.basename(file.filename or "图片")[:255]
|
||||
attachment = ChatAttachment(
|
||||
avatar_id=avatar.id,
|
||||
uploader_kind=uploader_kind,
|
||||
filename=filename,
|
||||
mime_type=(file.content_type or "")[:100],
|
||||
file_size=len(content),
|
||||
status="processing",
|
||||
expires_at=_attachment_expiry(),
|
||||
)
|
||||
db.add(attachment)
|
||||
db.commit()
|
||||
db.refresh(attachment)
|
||||
|
||||
try:
|
||||
prepared = prepare_image(content)
|
||||
model_config = get_chat_model_config()
|
||||
vision_result = _run_billed_vision_call(
|
||||
db,
|
||||
avatar,
|
||||
prepared,
|
||||
model=model_config.vision_model,
|
||||
prompt=GENERAL_VISION_PROMPT,
|
||||
source="vision_image",
|
||||
json_output=True,
|
||||
model_config=model_config,
|
||||
)
|
||||
analysis = parse_vision_analysis(vision_result["content"])
|
||||
extracted_text = analysis.get("visible_text") or ""
|
||||
ocr_model = ""
|
||||
ocr_failed = False
|
||||
if analysis["category"] == "medical_document" and model_config.ocr_model:
|
||||
try:
|
||||
ocr_result = _run_billed_vision_call(
|
||||
db,
|
||||
avatar,
|
||||
prepared,
|
||||
model=model_config.ocr_model,
|
||||
prompt=MEDICAL_OCR_PROMPT,
|
||||
source="vision_medical_ocr",
|
||||
json_output=False,
|
||||
model_config=model_config,
|
||||
)
|
||||
extracted_text = ocr_result["content"]
|
||||
ocr_model = model_config.ocr_model
|
||||
except (RuntimeError, InsufficientTokensError):
|
||||
ocr_failed = True
|
||||
logger.warning(
|
||||
"medical OCR degraded for attachment %s avatar %s",
|
||||
attachment.id,
|
||||
avatar.id,
|
||||
)
|
||||
|
||||
attachment.mime_type = prepared.mime_type
|
||||
attachment.status = "ready"
|
||||
attachment.category = analysis["category"]
|
||||
attachment.summary = analysis.get("summary") or "图片内容已识别"
|
||||
attachment.extracted_text = extracted_text
|
||||
attachment.structured_data = analysis
|
||||
attachment.warning = build_attachment_warning(analysis, ocr_failed=ocr_failed)
|
||||
attachment.vision_model = model_config.vision_model
|
||||
attachment.ocr_model = ocr_model
|
||||
db.commit()
|
||||
db.refresh(attachment)
|
||||
logger.info(
|
||||
"chat image ready attachment=%s avatar=%s category=%s model=%s ocr=%s",
|
||||
attachment.id,
|
||||
avatar.id,
|
||||
attachment.category,
|
||||
attachment.vision_model,
|
||||
bool(attachment.ocr_model),
|
||||
)
|
||||
return attachment
|
||||
except ImageValidationError as exc:
|
||||
attachment.status = "failed"
|
||||
attachment.warning = str(exc)
|
||||
db.commit()
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
except InsufficientTokensError:
|
||||
attachment.status = "failed"
|
||||
attachment.warning = "积分余额不足"
|
||||
db.commit()
|
||||
raise
|
||||
except RuntimeError as exc:
|
||||
attachment.status = "failed"
|
||||
attachment.warning = str(exc)
|
||||
db.commit()
|
||||
logger.warning(
|
||||
"chat image failed attachment=%s avatar=%s error=%s",
|
||||
attachment.id,
|
||||
avatar.id,
|
||||
type(exc).__name__,
|
||||
)
|
||||
raise HTTPException(status_code=502, detail=str(exc)) from exc
|
||||
finally:
|
||||
content = b""
|
||||
|
||||
|
||||
def _normalize_question(value: str) -> str:
|
||||
value = (value or "").strip().lower()
|
||||
value = re.sub(r"\s+", "", value)
|
||||
@@ -233,6 +489,7 @@ def _build_prompt(
|
||||
knowledge_hits: list[dict],
|
||||
*,
|
||||
standard_answer: str = "",
|
||||
image_contexts: list[dict] | None = None,
|
||||
) -> list[dict]:
|
||||
config = _config(avatar)
|
||||
description = (getattr(avatar, "description", "") or "").strip()
|
||||
@@ -241,6 +498,7 @@ def _build_prompt(
|
||||
for hit in knowledge_hits
|
||||
if hit.get("snippet")
|
||||
)
|
||||
image_contexts = image_contexts or []
|
||||
profile_items = [
|
||||
(label, config[key])
|
||||
for label, key in (
|
||||
@@ -265,6 +523,24 @@ def _build_prompt(
|
||||
)
|
||||
if config["systemPrompt"]:
|
||||
system += f"\n额外系统提示词:{config['systemPrompt']}"
|
||||
if image_contexts:
|
||||
image_material = json.dumps(image_contexts, ensure_ascii=False, default=str)
|
||||
system += (
|
||||
"\n以下是当前会话图片经过视觉识别后得到的资料:\n"
|
||||
f"{image_material}"
|
||||
"\n图片资料可能包含 OCR 错字、模糊内容或用户尚未确认的信息,只能按可见内容谨慎表达。"
|
||||
"标准答题对中的事实优先级高于图片资料,知识库事实优先级高于模型推测;发生冲突时遵循更高优先级资料,"
|
||||
"并自然提醒对方核对原图。不得声称看到了图片中不存在的内容。"
|
||||
)
|
||||
if any(
|
||||
context.get("category") in {"medical_document", "medical_image"}
|
||||
for context in image_contexts
|
||||
):
|
||||
system += (
|
||||
"\n本次包含医疗资料。可以整理病例原文、解释指标含义和提示需要关注的异常,但不能仅凭图片作出"
|
||||
"确定诊断、疾病分期、处方、停药或治疗决定。医学影像只能客观描述,并提醒结合正规报告和医生意见。"
|
||||
"回答结尾用与用户相同的语言简短说明图片识别结果仅供辅助,不能替代医生诊断。"
|
||||
)
|
||||
if standard_answer:
|
||||
system += (
|
||||
f"\n以下是本次问题命中的已确认标准答案:\n{standard_answer.strip()}"
|
||||
@@ -277,6 +553,11 @@ def _build_prompt(
|
||||
"\n涉及事实、专业判断、地址、流程、数据或建议时,只能依据本人资料、标准问答形成的上下文"
|
||||
"和以上可靠资料作答,不要补充资料之外的通用知识或自行推测。"
|
||||
)
|
||||
elif image_contexts:
|
||||
system += (
|
||||
"\n本次没有命中标准答题对或文件知识库,但已提供图片识别资料。只能围绕图片中的可确认内容、"
|
||||
"本人资料和当前对话作答;不要补充图片之外的事实、专业判断或具体建议。"
|
||||
)
|
||||
else:
|
||||
system += (
|
||||
"\n本次问题没有检索到可靠资料。除自然寒暄和基于本人资料的回答外,不要凭通用知识给出事实、"
|
||||
@@ -439,14 +720,17 @@ def _resolve_reply(
|
||||
search_fn: Callable[..., list[dict]] | None = None,
|
||||
model_client: Callable[..., str] | None = None,
|
||||
usage_source: str = "chat",
|
||||
image_contexts: list[dict] | None = None,
|
||||
) -> dict:
|
||||
image_contexts = image_contexts or []
|
||||
question = question.strip() or "请根据这张图片说明可确认的内容。"
|
||||
if qa_pairs is None:
|
||||
qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all()
|
||||
matched = _match_standard_qa(question, qa_pairs)
|
||||
adapt_qa_language = bool(
|
||||
matched and _qa_requires_language_adaptation(question, matched.answer)
|
||||
)
|
||||
if matched and not adapt_qa_language:
|
||||
if matched and not adapt_qa_language and not image_contexts:
|
||||
return {"answer": matched.answer, "source": "qa", "references": []}
|
||||
|
||||
if matched:
|
||||
@@ -457,11 +741,19 @@ def _resolve_reply(
|
||||
question,
|
||||
hits,
|
||||
standard_answer=matched.answer,
|
||||
image_contexts=image_contexts,
|
||||
)
|
||||
else:
|
||||
search_fn = search_fn or (lambda query, avatar_id: _search_knowledge(db, avatar_id, query))
|
||||
hits = search_fn(question, avatar.id)
|
||||
messages = _build_prompt(avatar, history, question, hits)
|
||||
retrieval_question = _image_retrieval_question(question, image_contexts)
|
||||
hits = search_fn(retrieval_question, avatar.id)
|
||||
messages = _build_prompt(
|
||||
avatar,
|
||||
history,
|
||||
question,
|
||||
hits,
|
||||
image_contexts=image_contexts,
|
||||
)
|
||||
config = _config(avatar)
|
||||
temperature = 0.0 if matched else min(
|
||||
0.45 if hits else 0.25,
|
||||
@@ -498,7 +790,9 @@ def _resolve_reply(
|
||||
raise
|
||||
result = {
|
||||
"answer": answer,
|
||||
"source": "qa" if matched else ("knowledge" if hits else "qwen"),
|
||||
"source": "qa" if matched else (
|
||||
"knowledge" if hits else ("vision" if image_contexts else "qwen")
|
||||
),
|
||||
"references": hits,
|
||||
}
|
||||
if token_usage:
|
||||
@@ -514,14 +808,17 @@ def _stream_reply(
|
||||
*,
|
||||
public: bool = False,
|
||||
usage_source: str = "chat_stream",
|
||||
image_contexts: list[dict] | None = None,
|
||||
):
|
||||
image_contexts = image_contexts or []
|
||||
question = question.strip() or "请根据这张图片说明可确认的内容。"
|
||||
qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all()
|
||||
matched = _match_standard_qa(question, qa_pairs)
|
||||
adapt_qa_language = bool(
|
||||
matched and _qa_requires_language_adaptation(question, matched.answer)
|
||||
)
|
||||
messages, reservation = [], None
|
||||
if matched and not adapt_qa_language:
|
||||
if matched and not adapt_qa_language and not image_contexts:
|
||||
source, references, chunks = "qa", [], _iter_text_chunks(matched.answer)
|
||||
else:
|
||||
if matched:
|
||||
@@ -533,11 +830,21 @@ def _stream_reply(
|
||||
question,
|
||||
references,
|
||||
standard_answer=matched.answer,
|
||||
image_contexts=image_contexts,
|
||||
)
|
||||
else:
|
||||
references = _search_knowledge(db, avatar.id, question)
|
||||
source = "knowledge" if references else "qwen"
|
||||
messages = _build_prompt(avatar, history, question, references)
|
||||
retrieval_question = _image_retrieval_question(question, image_contexts)
|
||||
references = _search_knowledge(db, avatar.id, retrieval_question)
|
||||
source = "knowledge" if references else (
|
||||
"vision" if image_contexts else "qwen"
|
||||
)
|
||||
messages = _build_prompt(
|
||||
avatar,
|
||||
history,
|
||||
question,
|
||||
references,
|
||||
image_contexts=image_contexts,
|
||||
)
|
||||
config = _config(avatar)
|
||||
temperature = 0.0 if matched else min(
|
||||
0.45 if references else 0.25,
|
||||
@@ -641,11 +948,62 @@ def get_shared_avatar(share_token: str, db: Session = Depends(get_db)):
|
||||
return ok(_public_avatar_payload(_require_shared_avatar(db, share_token)))
|
||||
|
||||
|
||||
@router.post("/avatar/{avatar_id}/chat/images")
|
||||
async def upload_chat_image(
|
||||
avatar_id: str,
|
||||
file: UploadFile = File(...),
|
||||
authorization: str = Header(None),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
avatar = _require_owned_avatar(db, avatar_id, authorization)
|
||||
purge_expired_chat_attachments(db)
|
||||
try:
|
||||
attachment = await _analyze_uploaded_image(
|
||||
db,
|
||||
avatar,
|
||||
file,
|
||||
uploader_kind="owner",
|
||||
)
|
||||
return ok(attachment.to_dict())
|
||||
except InsufficientTokensError as exc:
|
||||
raise HTTPException(status_code=402, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.post("/public/avatar/{share_token}/chat/images")
|
||||
async def upload_public_chat_image(
|
||||
share_token: str,
|
||||
file: UploadFile = File(...),
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
avatar = _require_shared_avatar(db, share_token)
|
||||
purge_expired_chat_attachments(db)
|
||||
try:
|
||||
attachment = await _analyze_uploaded_image(
|
||||
db,
|
||||
avatar,
|
||||
file,
|
||||
uploader_kind="public",
|
||||
)
|
||||
return ok(attachment.to_dict())
|
||||
except InsufficientTokensError as exc:
|
||||
raise HTTPException(status_code=402, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.post("/public/avatar/{share_token}/chat")
|
||||
def public_chat(share_token: str, body: ChatIn = Body(...), db: Session = Depends(get_db)):
|
||||
avatar = _require_shared_avatar(db, share_token)
|
||||
image_contexts = _attachment_contexts(
|
||||
_load_chat_attachments(db, avatar.id, body)
|
||||
)
|
||||
try:
|
||||
result = _resolve_reply(db, avatar, body.message, body.history, usage_source="public_chat")
|
||||
result = _resolve_reply(
|
||||
db,
|
||||
avatar,
|
||||
body.message,
|
||||
body.history,
|
||||
usage_source="public_chat",
|
||||
image_contexts=image_contexts,
|
||||
)
|
||||
# 公开访客无需获知知识文件名、检索分数或内部答复来源。
|
||||
result["references"] = []
|
||||
result["source"] = "public"
|
||||
@@ -660,8 +1018,17 @@ def public_chat(share_token: str, body: ChatIn = Body(...), db: Session = Depend
|
||||
@router.post("/avatar/{avatar_id}/chat")
|
||||
def chat(avatar_id: str, body: ChatIn = Body(...), authorization: str = Header(None), db: Session = Depends(get_db)):
|
||||
avatar = _require_owned_avatar(db, avatar_id, authorization)
|
||||
image_contexts = _attachment_contexts(
|
||||
_load_chat_attachments(db, avatar.id, body)
|
||||
)
|
||||
try:
|
||||
return ok(_resolve_reply(db, avatar, body.message, body.history))
|
||||
return ok(_resolve_reply(
|
||||
db,
|
||||
avatar,
|
||||
body.message,
|
||||
body.history,
|
||||
image_contexts=image_contexts,
|
||||
))
|
||||
except InsufficientTokensError as exc:
|
||||
return fail(str(exc), code=402)
|
||||
except RuntimeError as exc:
|
||||
@@ -671,7 +1038,17 @@ def chat(avatar_id: str, body: ChatIn = Body(...), authorization: str = Header(N
|
||||
@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)):
|
||||
try:
|
||||
return _stream_reply(db, _require_owned_avatar(db, avatar_id, authorization), body.message, body.history)
|
||||
avatar = _require_owned_avatar(db, avatar_id, authorization)
|
||||
image_contexts = _attachment_contexts(
|
||||
_load_chat_attachments(db, avatar.id, body)
|
||||
)
|
||||
return _stream_reply(
|
||||
db,
|
||||
avatar,
|
||||
body.message,
|
||||
body.history,
|
||||
image_contexts=image_contexts,
|
||||
)
|
||||
except InsufficientTokensError as exc:
|
||||
raise HTTPException(status_code=402, detail=str(exc)) from exc
|
||||
|
||||
@@ -679,13 +1056,18 @@ def chat_stream(avatar_id: str, body: ChatIn = Body(...), authorization: str = H
|
||||
@router.post("/public/avatar/{share_token}/chat/stream")
|
||||
def public_chat_stream(share_token: str, body: ChatIn = Body(...), db: Session = Depends(get_db)):
|
||||
try:
|
||||
avatar = _require_shared_avatar(db, share_token)
|
||||
image_contexts = _attachment_contexts(
|
||||
_load_chat_attachments(db, avatar.id, body)
|
||||
)
|
||||
return _stream_reply(
|
||||
db,
|
||||
_require_shared_avatar(db, share_token),
|
||||
avatar,
|
||||
body.message,
|
||||
body.history,
|
||||
public=True,
|
||||
usage_source="public_chat_stream",
|
||||
image_contexts=image_contexts,
|
||||
)
|
||||
except InsufficientTokensError as exc:
|
||||
raise HTTPException(status_code=402, detail=str(exc)) from exc
|
||||
|
||||
Reference in New Issue
Block a user