feat: isolate external application conversations
This commit is contained in:
@@ -18,6 +18,7 @@ from app.models.admin import Admin
|
||||
from app.models.chat import ChatMessage, ChatSession, TopicSession
|
||||
from app.models.growth import ShareDraft, TeacherHelpCard, TopicSummary
|
||||
from app.models.logs import AiRequestLog, OperationLog
|
||||
from app.models.sso import SsoClient
|
||||
from app.models.user import User
|
||||
from app.api.pagination import page_result
|
||||
from app.services.admin_service import OperationLogService
|
||||
@@ -34,6 +35,8 @@ def chat_list(
|
||||
keyword: str = Query(default=""),
|
||||
userId: int | None = Query(default=None),
|
||||
status: str = Query(default=""),
|
||||
sourceType: str = Query(default="", pattern="^(|direct|sso)$"),
|
||||
sourceClientId: int | None = Query(default=None, gt=0),
|
||||
dateFrom: datetime | None = Query(default=None),
|
||||
dateTo: datetime | None = Query(default=None),
|
||||
page: int = Query(default=1, ge=1),
|
||||
@@ -41,10 +44,19 @@ def chat_list(
|
||||
db: Session = Depends(get_db),
|
||||
current_admin: Admin = Depends(get_current_admin),
|
||||
) -> dict:
|
||||
query = _chat_query(keyword=keyword, user_id=userId, status=status, date_from=dateFrom, date_to=dateTo)
|
||||
query = _chat_query(
|
||||
keyword=keyword,
|
||||
user_id=userId,
|
||||
status=status,
|
||||
source_type=sourceType,
|
||||
source_client_id=sourceClientId,
|
||||
date_from=dateFrom,
|
||||
date_to=dateTo,
|
||||
)
|
||||
total = db.scalar(select(func.count()).select_from(query.order_by(None).subquery())) or 0
|
||||
rows = db.execute(query.order_by(ChatSession.updated_at.desc()).offset((page - 1) * pageSize).limit(pageSize)).all()
|
||||
return api_success(page_result([_chat_row_dict(session, user) for session, user in rows], total=total, page=page, page_size=pageSize))
|
||||
items = [_chat_row_dict(session, user, source_client) for session, user, source_client in rows]
|
||||
return api_success(page_result(items, total=total, page=page, page_size=pageSize))
|
||||
|
||||
|
||||
@router.get("/chat/export")
|
||||
@@ -52,26 +64,37 @@ def export_chats(
|
||||
keyword: str = Query(default=""),
|
||||
userId: int | None = Query(default=None),
|
||||
status: str = Query(default=""),
|
||||
sourceType: str = Query(default="", pattern="^(|direct|sso)$"),
|
||||
sourceClientId: int | None = Query(default=None, gt=0),
|
||||
dateFrom: datetime | None = Query(default=None),
|
||||
dateTo: datetime | None = Query(default=None),
|
||||
db: Session = Depends(get_db),
|
||||
current_admin: Admin = Depends(get_current_admin),
|
||||
) -> Response:
|
||||
rows = db.execute(
|
||||
_chat_query(keyword=keyword, user_id=userId, status=status, date_from=dateFrom, date_to=dateTo)
|
||||
_chat_query(
|
||||
keyword=keyword,
|
||||
user_id=userId,
|
||||
status=status,
|
||||
source_type=sourceType,
|
||||
source_client_id=sourceClientId,
|
||||
date_from=dateFrom,
|
||||
date_to=dateTo,
|
||||
)
|
||||
.order_by(ChatSession.updated_at.desc())
|
||||
.limit(1000)
|
||||
).all()
|
||||
output = StringIO()
|
||||
writer = csv.writer(output)
|
||||
writer.writerow(["会话ID", "用户ID", "手机号", "姓名", "标题", "消息数", "最后消息时间", "更新时间"])
|
||||
for session, user in rows:
|
||||
writer.writerow(["会话ID", "用户ID", "手机号", "姓名", "来源", "标题", "消息数", "最后消息时间", "更新时间"])
|
||||
for session, user, source_client in rows:
|
||||
writer.writerow(
|
||||
[
|
||||
session.id,
|
||||
session.user_id,
|
||||
user.phone if user else "",
|
||||
user.name if user else "",
|
||||
source_client.name if source_client else "千问千答直接访问",
|
||||
session.title,
|
||||
session.message_count,
|
||||
session.last_message_at,
|
||||
@@ -117,8 +140,9 @@ def chat_detail(
|
||||
current_admin: Admin = Depends(get_current_admin),
|
||||
) -> dict:
|
||||
session_row = db.execute(
|
||||
select(ChatSession, User)
|
||||
select(ChatSession, User, SsoClient)
|
||||
.join(User, User.id == ChatSession.user_id, isouter=True)
|
||||
.join(SsoClient, SsoClient.id == ChatSession.source_client_id, isouter=True)
|
||||
.where(ChatSession.id == session_id)
|
||||
).first()
|
||||
if session_row is None:
|
||||
@@ -126,7 +150,7 @@ def chat_detail(
|
||||
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="会话不存在")
|
||||
|
||||
session, user = session_row
|
||||
session, user, source_client = session_row
|
||||
message_query = select(ChatMessage).where(ChatMessage.session_id == session_id)
|
||||
message_total = db.scalar(select(func.count()).select_from(message_query.subquery())) or 0
|
||||
messages = db.scalars(
|
||||
@@ -161,7 +185,7 @@ def chat_detail(
|
||||
)
|
||||
return api_success(
|
||||
{
|
||||
"session": _chat_row_dict(session, user),
|
||||
"session": _chat_row_dict(session, user, source_client),
|
||||
"messages": [_message_dict(item) for item in messages],
|
||||
"messagesPage": page_result(
|
||||
[_message_dict(item) for item in messages],
|
||||
@@ -305,16 +329,23 @@ def _chat_query(
|
||||
keyword: str,
|
||||
user_id: int | None,
|
||||
status: str,
|
||||
source_type: str,
|
||||
source_client_id: int | None,
|
||||
date_from: datetime | None,
|
||||
date_to: datetime | None,
|
||||
):
|
||||
query = (
|
||||
select(ChatSession, User)
|
||||
select(ChatSession, User, SsoClient)
|
||||
.join(User, User.id == ChatSession.user_id, isouter=True)
|
||||
.join(SsoClient, SsoClient.id == ChatSession.source_client_id, isouter=True)
|
||||
.where(ChatSession.is_deleted == 0)
|
||||
)
|
||||
if user_id is not None:
|
||||
query = query.where(ChatSession.user_id == user_id)
|
||||
if source_type:
|
||||
query = query.where(ChatSession.source_type == source_type)
|
||||
if source_client_id is not None:
|
||||
query = query.where(ChatSession.source_type == "sso", ChatSession.source_client_id == source_client_id)
|
||||
if date_from is not None:
|
||||
query = query.where(ChatSession.updated_at >= date_from.replace(tzinfo=None))
|
||||
if date_to is not None:
|
||||
@@ -340,12 +371,15 @@ def _chat_query(
|
||||
return query
|
||||
|
||||
|
||||
def _chat_row_dict(session: ChatSession, user: User | None) -> dict:
|
||||
def _chat_row_dict(session: ChatSession, user: User | None, source_client: SsoClient | None = None) -> dict:
|
||||
return {
|
||||
"id": session.id,
|
||||
"userId": session.user_id,
|
||||
"userPhone": user.phone if user else "",
|
||||
"userName": user.name if user else "",
|
||||
"sourceType": session.source_type,
|
||||
"sourceClientId": session.source_client_id,
|
||||
"sourceName": source_client.name if source_client else "千问千答直接访问",
|
||||
"title": session.title,
|
||||
"messageCount": session.message_count,
|
||||
"lastMessageAt": session.last_message_at,
|
||||
|
||||
@@ -13,6 +13,7 @@ from app.core.dependencies import get_current_admin
|
||||
from app.core.responses import api_success
|
||||
from app.models.admin import Admin
|
||||
from app.models.ai_config import SystemConfig
|
||||
from app.models.entitlement import EntitlementPlan
|
||||
from app.models.sso import SsoClient, SsoLoginAudit, UserExternalIdentity
|
||||
from app.models.user import User
|
||||
from app.schemas.sso import SsoClientSaveRequest, SsoClientUpdateRequest, SsoPublicConfigRequest
|
||||
@@ -80,11 +81,12 @@ def list_sso_clients(
|
||||
.subquery()
|
||||
)
|
||||
rows = db.execute(
|
||||
select(SsoClient, func.coalesce(identity_counts.c.identity_count, 0))
|
||||
select(SsoClient, func.coalesce(identity_counts.c.identity_count, 0), EntitlementPlan.name)
|
||||
.outerjoin(identity_counts, identity_counts.c.client_id == SsoClient.id)
|
||||
.outerjoin(EntitlementPlan, EntitlementPlan.id == SsoClient.default_entitlement_plan_id)
|
||||
.order_by(SsoClient.created_at.desc(), SsoClient.id.desc())
|
||||
).all()
|
||||
return api_success([_client_item(client, int(identity_count)) for client, identity_count in rows])
|
||||
return api_success([_client_item(client, int(identity_count), plan_name) for client, identity_count, plan_name in rows])
|
||||
|
||||
|
||||
@router.post("/sso/client")
|
||||
@@ -96,12 +98,15 @@ def create_sso_client(
|
||||
app_id = payload.appId.strip()
|
||||
if db.scalar(select(SsoClient.id).where(SsoClient.app_id == app_id)) is not None:
|
||||
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="应用ID已存在")
|
||||
_validate_default_plan(db, payload.allowAutoRegister, payload.defaultEntitlementPlanId)
|
||||
plaintext_secret = secrets.token_urlsafe(32)
|
||||
client = SsoClient(
|
||||
app_id=app_id,
|
||||
name=payload.name.strip(),
|
||||
client_secret=SecretService.encrypt(plaintext_secret),
|
||||
redirect_uris=json.dumps(payload.redirectUris, ensure_ascii=False),
|
||||
allow_auto_register=1 if payload.allowAutoRegister else 0,
|
||||
default_entitlement_plan_id=payload.defaultEntitlementPlanId,
|
||||
status=payload.status,
|
||||
created_by=current_admin.id,
|
||||
)
|
||||
@@ -127,8 +132,11 @@ def update_sso_client(
|
||||
current_admin: Admin = Depends(get_current_admin),
|
||||
) -> dict:
|
||||
client = _require_client(db, client_id)
|
||||
_validate_default_plan(db, payload.allowAutoRegister, payload.defaultEntitlementPlanId)
|
||||
client.name = payload.name.strip()
|
||||
client.redirect_uris = json.dumps(payload.redirectUris, ensure_ascii=False)
|
||||
client.allow_auto_register = 1 if payload.allowAutoRegister else 0
|
||||
client.default_entitlement_plan_id = payload.defaultEntitlementPlanId
|
||||
client.status = payload.status
|
||||
OperationLogService.write(
|
||||
db,
|
||||
@@ -141,7 +149,8 @@ def update_sso_client(
|
||||
identity_count = db.scalar(
|
||||
select(func.count(UserExternalIdentity.id)).where(UserExternalIdentity.client_id == client.id)
|
||||
) or 0
|
||||
return api_success(_client_item(client, identity_count))
|
||||
plan_name = db.scalar(select(EntitlementPlan.name).where(EntitlementPlan.id == client.default_entitlement_plan_id))
|
||||
return api_success(_client_item(client, identity_count, plan_name))
|
||||
|
||||
|
||||
@router.post("/sso/client/{client_id}/secret/rotate")
|
||||
@@ -260,12 +269,15 @@ def _require_client(db: Session, client_id: int) -> SsoClient:
|
||||
return client
|
||||
|
||||
|
||||
def _client_item(client: SsoClient, identity_count: int) -> dict:
|
||||
def _client_item(client: SsoClient, identity_count: int, plan_name: str | None = None) -> dict:
|
||||
return {
|
||||
"id": client.id,
|
||||
"appId": client.app_id,
|
||||
"name": client.name,
|
||||
"redirectUris": _json_list(client.redirect_uris),
|
||||
"allowAutoRegister": bool(client.allow_auto_register),
|
||||
"defaultEntitlementPlanId": client.default_entitlement_plan_id,
|
||||
"defaultEntitlementPlanName": plan_name,
|
||||
"status": client.status,
|
||||
"identityCount": identity_count,
|
||||
"lastUsedAt": client.last_used_at,
|
||||
@@ -274,6 +286,14 @@ def _client_item(client: SsoClient, identity_count: int) -> dict:
|
||||
}
|
||||
|
||||
|
||||
def _validate_default_plan(db: Session, allow_auto_register: bool, plan_id: int | None) -> None:
|
||||
if not allow_auto_register:
|
||||
return
|
||||
plan = db.get(EntitlementPlan, plan_id) if plan_id else None
|
||||
if plan is None or plan.status != 1 or plan.plan_type == "teacher":
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="默认权益版本不存在或已停用")
|
||||
|
||||
|
||||
def _identity_item(identity: UserExternalIdentity, client: SsoClient, user: User) -> dict:
|
||||
return {
|
||||
"id": identity.id,
|
||||
|
||||
@@ -490,6 +490,8 @@ def _user_dict(user: User, entitlement: dict | None = None) -> dict:
|
||||
"phone": user.phone,
|
||||
"name": user.name,
|
||||
"nickname": user.nickname,
|
||||
"registrationSource": user.registration_source,
|
||||
"registrationClientId": user.registration_client_id,
|
||||
"status": user.status,
|
||||
"dailyChatLimit": user.daily_chat_limit,
|
||||
"dailyChatUsed": user.daily_chat_used,
|
||||
|
||||
@@ -7,12 +7,15 @@ from collections.abc import AsyncIterator
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from fastapi.responses import StreamingResponse
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.core.database import get_db
|
||||
from app.core.dependencies import get_current_user
|
||||
from app.core.auth_context import ChatAccessScope, UserAuthContext
|
||||
from app.core.dependencies import get_current_user_context
|
||||
from app.core.responses import api_success
|
||||
from app.models.user import User
|
||||
from app.models.chat import TopicSession
|
||||
from app.models.growth import TopicSummary
|
||||
from app.schemas.chat import (
|
||||
ChatCompletionRequest,
|
||||
ChatMessageRead,
|
||||
@@ -40,14 +43,14 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@router.post("/session")
|
||||
def create_session(db: Session = Depends(get_db), current_user: User = Depends(get_current_user)) -> dict:
|
||||
session = ChatService.create_session(db, current_user)
|
||||
def create_session(db: Session = Depends(get_db), current: UserAuthContext = Depends(get_current_user_context)) -> dict:
|
||||
session = ChatService.create_session(db, current.user, current.chat_scope)
|
||||
return api_success(CreateSessionResponse(sessionId=session.id).model_dump())
|
||||
|
||||
|
||||
@router.get("/session/list")
|
||||
def list_sessions(db: Session = Depends(get_db), current_user: User = Depends(get_current_user)) -> dict:
|
||||
sessions = ChatService.list_sessions(db, current_user)
|
||||
def list_sessions(db: Session = Depends(get_db), current: UserAuthContext = Depends(get_current_user_context)) -> dict:
|
||||
sessions = ChatService.list_sessions(db, current.user, current.chat_scope)
|
||||
return api_success([ChatSessionRead.model_validate(session).model_dump() for session in sessions])
|
||||
|
||||
|
||||
@@ -55,9 +58,9 @@ def list_sessions(db: Session = Depends(get_db), current_user: User = Depends(ge
|
||||
def history(
|
||||
sessionId: int = Query(...),
|
||||
db: Session = Depends(get_db),
|
||||
current_user: User = Depends(get_current_user),
|
||||
current: UserAuthContext = Depends(get_current_user_context),
|
||||
) -> dict:
|
||||
messages = ChatService.get_history(db, current_user, sessionId)
|
||||
messages = ChatService.get_history(db, current.user, sessionId, current.chat_scope)
|
||||
reasoning_visible = ReasoningPolicyService.is_visible(db)
|
||||
result = []
|
||||
for message in messages:
|
||||
@@ -72,9 +75,9 @@ def history(
|
||||
def update_title(
|
||||
payload: UpdateSessionTitleRequest,
|
||||
db: Session = Depends(get_db),
|
||||
current_user: User = Depends(get_current_user),
|
||||
current: UserAuthContext = Depends(get_current_user_context),
|
||||
) -> dict:
|
||||
session = ChatService.update_title(db, current_user, payload.sessionId, payload.title)
|
||||
session = ChatService.update_title(db, current.user, payload.sessionId, payload.title, current.chat_scope)
|
||||
return api_success(ChatSessionRead.model_validate(session).model_dump())
|
||||
|
||||
|
||||
@@ -82,9 +85,9 @@ def update_title(
|
||||
def delete_session(
|
||||
session_id: int,
|
||||
db: Session = Depends(get_db),
|
||||
current_user: User = Depends(get_current_user),
|
||||
current: UserAuthContext = Depends(get_current_user_context),
|
||||
) -> dict:
|
||||
ChatService.delete_session(db, current_user, session_id)
|
||||
ChatService.delete_session(db, current.user, session_id, current.chat_scope)
|
||||
return api_success()
|
||||
|
||||
|
||||
@@ -92,29 +95,37 @@ def delete_session(
|
||||
def finish_topic(
|
||||
session_id: int,
|
||||
db: Session = Depends(get_db),
|
||||
current_user: User = Depends(get_current_user),
|
||||
current: UserAuthContext = Depends(get_current_user_context),
|
||||
) -> dict:
|
||||
session = ChatService._get_user_session(db, current_user, session_id)
|
||||
return api_success(GrowthProfileService.finish_active_topic(db, user=current_user, session=session))
|
||||
session = ChatService._get_user_session(db, current.user, session_id, current.chat_scope)
|
||||
return api_success(GrowthProfileService.finish_active_topic(db, user=current.user, session=session))
|
||||
|
||||
|
||||
@router.get("/topic/settlement/{summary_id}")
|
||||
def topic_settlement(
|
||||
summary_id: int,
|
||||
db: Session = Depends(get_db),
|
||||
current_user: User = Depends(get_current_user),
|
||||
current: UserAuthContext = Depends(get_current_user_context),
|
||||
) -> dict:
|
||||
return api_success(GrowthProfileService.topic_settlement_result(db, user=current_user, summary_id=summary_id))
|
||||
topic_session_id = db.scalar(
|
||||
select(TopicSession.chat_session_id)
|
||||
.join(TopicSummary, TopicSummary.topic_session_id == TopicSession.id)
|
||||
.where(TopicSummary.id == summary_id, TopicSummary.user_id == current.user.id)
|
||||
)
|
||||
if topic_session_id is None:
|
||||
raise HTTPException(status_code=404, detail="主题沉淀任务不存在")
|
||||
ChatService._get_user_session(db, current.user, topic_session_id, current.chat_scope)
|
||||
return api_success(GrowthProfileService.topic_settlement_result(db, user=current.user, summary_id=summary_id))
|
||||
|
||||
|
||||
@router.post("/session/{session_id}/help-card")
|
||||
def generate_help_card(
|
||||
session_id: int,
|
||||
db: Session = Depends(get_db),
|
||||
current_user: User = Depends(get_current_user),
|
||||
current: UserAuthContext = Depends(get_current_user_context),
|
||||
) -> dict:
|
||||
session = ChatService._get_user_session(db, current_user, session_id)
|
||||
card = HelpCardService.generate_for_session(db, user=current_user, session=session)
|
||||
session = ChatService._get_user_session(db, current.user, session_id, current.chat_scope)
|
||||
card = HelpCardService.generate_for_session(db, user=current.user, session=session)
|
||||
return api_success(help_card_dict(card))
|
||||
|
||||
|
||||
@@ -122,27 +133,33 @@ def generate_help_card(
|
||||
def list_help_cards(
|
||||
limit: int = Query(default=20, ge=1, le=50),
|
||||
db: Session = Depends(get_db),
|
||||
current_user: User = Depends(get_current_user),
|
||||
current: UserAuthContext = Depends(get_current_user_context),
|
||||
) -> dict:
|
||||
return api_success([help_card_dict(card) for card in HelpCardService.list_user_cards(db, user=current_user, limit=limit)])
|
||||
cards = HelpCardService.list_user_cards(
|
||||
db,
|
||||
user=current.user,
|
||||
scope=current.chat_scope,
|
||||
limit=limit,
|
||||
)
|
||||
return api_success([help_card_dict(card) for card in cards])
|
||||
|
||||
|
||||
@router.post("/help-card/{card_id}/copied")
|
||||
def mark_help_card_copied(
|
||||
card_id: int,
|
||||
db: Session = Depends(get_db),
|
||||
current_user: User = Depends(get_current_user),
|
||||
current: UserAuthContext = Depends(get_current_user_context),
|
||||
) -> dict:
|
||||
return api_success(help_card_dict(HelpCardService.mark_copied(db, user=current_user, card_id=card_id)))
|
||||
return api_success(help_card_dict(HelpCardService.mark_copied(db, user=current.user, scope=current.chat_scope, card_id=card_id)))
|
||||
|
||||
|
||||
@router.delete("/help-card/{card_id}")
|
||||
def delete_help_card(
|
||||
card_id: int,
|
||||
db: Session = Depends(get_db),
|
||||
current_user: User = Depends(get_current_user),
|
||||
current: UserAuthContext = Depends(get_current_user_context),
|
||||
) -> dict:
|
||||
HelpCardService.delete(db, user=current_user, card_id=card_id)
|
||||
HelpCardService.delete(db, user=current.user, scope=current.chat_scope, card_id=card_id)
|
||||
return api_success()
|
||||
|
||||
|
||||
@@ -150,10 +167,10 @@ def delete_help_card(
|
||||
def generate_share_draft(
|
||||
session_id: int,
|
||||
db: Session = Depends(get_db),
|
||||
current_user: User = Depends(get_current_user),
|
||||
current: UserAuthContext = Depends(get_current_user_context),
|
||||
) -> dict:
|
||||
session = ChatService._get_user_session(db, current_user, session_id)
|
||||
draft = ShareDraftService.generate_for_session(db, user=current_user, session=session)
|
||||
session = ChatService._get_user_session(db, current.user, session_id, current.chat_scope)
|
||||
draft = ShareDraftService.generate_for_session(db, user=current.user, session=session)
|
||||
return api_success(share_draft_dict(draft))
|
||||
|
||||
|
||||
@@ -161,27 +178,33 @@ def generate_share_draft(
|
||||
def list_share_drafts(
|
||||
limit: int = Query(default=20, ge=1, le=50),
|
||||
db: Session = Depends(get_db),
|
||||
current_user: User = Depends(get_current_user),
|
||||
current: UserAuthContext = Depends(get_current_user_context),
|
||||
) -> dict:
|
||||
return api_success([share_draft_dict(draft) for draft in ShareDraftService.list_user_drafts(db, user=current_user, limit=limit)])
|
||||
drafts = ShareDraftService.list_user_drafts(
|
||||
db,
|
||||
user=current.user,
|
||||
scope=current.chat_scope,
|
||||
limit=limit,
|
||||
)
|
||||
return api_success([share_draft_dict(draft) for draft in drafts])
|
||||
|
||||
|
||||
@router.post("/share-draft/{draft_id}/copied")
|
||||
def mark_share_draft_copied(
|
||||
draft_id: int,
|
||||
db: Session = Depends(get_db),
|
||||
current_user: User = Depends(get_current_user),
|
||||
current: UserAuthContext = Depends(get_current_user_context),
|
||||
) -> dict:
|
||||
return api_success(share_draft_dict(ShareDraftService.mark_copied(db, user=current_user, draft_id=draft_id)))
|
||||
return api_success(share_draft_dict(ShareDraftService.mark_copied(db, user=current.user, scope=current.chat_scope, draft_id=draft_id)))
|
||||
|
||||
|
||||
@router.delete("/share-draft/{draft_id}")
|
||||
def delete_share_draft(
|
||||
draft_id: int,
|
||||
db: Session = Depends(get_db),
|
||||
current_user: User = Depends(get_current_user),
|
||||
current: UserAuthContext = Depends(get_current_user_context),
|
||||
) -> dict:
|
||||
ShareDraftService.delete(db, user=current_user, draft_id=draft_id)
|
||||
ShareDraftService.delete(db, user=current.user, scope=current.chat_scope, draft_id=draft_id)
|
||||
return api_success()
|
||||
|
||||
|
||||
@@ -189,10 +212,10 @@ def delete_share_draft(
|
||||
def completions(
|
||||
payload: ChatCompletionRequest,
|
||||
db: Session = Depends(get_db),
|
||||
current_user: User = Depends(get_current_user),
|
||||
current: UserAuthContext = Depends(get_current_user_context),
|
||||
) -> StreamingResponse:
|
||||
return StreamingResponse(
|
||||
_chat_stream(payload, db, current_user),
|
||||
_chat_stream(payload, db, current),
|
||||
media_type="text/event-stream",
|
||||
headers={
|
||||
"Cache-Control": "no-cache, no-transform",
|
||||
@@ -206,13 +229,17 @@ def completions(
|
||||
def stop(
|
||||
payload: StopChatRequest,
|
||||
db: Session = Depends(get_db),
|
||||
current_user: User = Depends(get_current_user),
|
||||
current: UserAuthContext = Depends(get_current_user_context),
|
||||
) -> dict:
|
||||
ChatService.stop_generation(db, current_user, payload.sessionId)
|
||||
ChatService.stop_generation(db, current.user, payload.sessionId, current.chat_scope)
|
||||
return api_success()
|
||||
|
||||
|
||||
async def _chat_stream(payload: ChatCompletionRequest, db: Session, current_user: User) -> AsyncIterator[str]:
|
||||
async def _chat_stream(payload: ChatCompletionRequest, db: Session, current: UserAuthContext) -> AsyncIterator[str]:
|
||||
# Keep the internal stream helper compatible with direct service/test callers;
|
||||
# HTTP requests always pass a fully validated UserAuthContext.
|
||||
current_user = getattr(current, "user", current)
|
||||
chat_scope = getattr(current, "chat_scope", ChatAccessScope.direct())
|
||||
config = load_chat_queue_config(db)
|
||||
queue_lease = await request_chat_slot(config)
|
||||
queue_request = queue_lease.request
|
||||
@@ -274,6 +301,7 @@ async def _chat_stream(payload: ChatCompletionRequest, db: Session, current_user
|
||||
payload.sessionId,
|
||||
payload.message,
|
||||
retry_failed_question=payload.retry,
|
||||
scope=chat_scope,
|
||||
)
|
||||
async for segment in ReasoningPolicyService.iter_segments(chunks):
|
||||
if segment.kind == "content":
|
||||
|
||||
Reference in New Issue
Block a user