from __future__ import annotations import json from collections.abc import AsyncGenerator from sqlalchemy import select from yuxi.channel.domain.model.message.stream_chat_request import StreamChatRequest from yuxi.repositories.agent_config_repository import AgentConfigRepository from yuxi.storage.postgres.models_business import User from yuxi.utils.logging_config import logger _CHANNEL_SERVICE_ROLE = "channel_service" class AgentAdapter: def __init__(self, session_factory): self._session_factory = session_factory async def _get_service_user(self, db): stmt = select(User).where(User.role == _CHANNEL_SERVICE_ROLE).limit(1) result = await db.execute(stmt) user = result.scalar_one_or_none() if user: return user stmt = select(User).where(User.role == "superadmin").limit(1) result = await db.execute(stmt) fallback = result.scalar_one_or_none() if fallback: logger.warning( "no %s user found, falling back to superadmin; create a %s user for production", _CHANNEL_SERVICE_ROLE, _CHANNEL_SERVICE_ROLE, ) return fallback async def stream_chat(self, request: StreamChatRequest) -> str: chunks = [] async for chunk in self.stream_chat_iter(request): chunks.append(chunk) return "".join(chunks) async def stream_chat_iter(self, request: StreamChatRequest) -> AsyncGenerator[str, None]: from yuxi.services.chat_service import stream_agent_chat async with self._session_factory() as db: service_user = await self._get_service_user(db) if not service_user: logger.error("no service user found for channel agent request") yield "Error: no service user found" return accumulated: list[str] = [] async for raw_chunk in stream_agent_chat( query=request.content, agent_config_id=request.agent_config_id, thread_id=request.session_id, meta={"request_id": request.message_id}, image_content=None, current_user=service_user, db=db, ): try: data = json.loads(raw_chunk) content = data.get("response") if content: accumulated.append(content) yield content status = data.get("status") if status == "error": error_msg = data.get("error_message", "unknown error") logger.error("stream_agent_chat error: %s", error_msg) if not accumulated: yield f"Error: {error_msg}" return except (json.JSONDecodeError, AttributeError): pass async def resolve_agent_id(self, agent_config_id: int) -> str: async with self._session_factory() as session: repo = AgentConfigRepository(session) config = await repo.get_by_id(config_id=agent_config_id) return config.agent_id if config else "chatbot"