import os import json import math import asyncio from uuid import UUID from datetime import date, datetime, timedelta, UTC from fastapi import HTTPException from tortoise.expressions import Q from langchain_openai import ChatOpenAI from langchain_anthropic import ChatAnthropic from langchain_google_genai import ChatGoogleGenerativeAI from langgraph.graph import START, END, MessagesState, StateGraph from langchain_core.tools import tool from langchain_community.tools import DuckDuckGoSearchResults from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver from langchain_community.document_loaders.blob_loaders import Blob from langchain_community.document_loaders.parsers import OpenAIWhisperParser from langchain_community.document_loaders import ( PyPDFLoader, Docx2txtLoader, UnstructuredExcelLoader, TextLoader, ) from langchain_core.messages import ( HumanMessage, AIMessage, SystemMessage, trim_messages, RemoveMessage, ToolMessage, ) from .database.service import DatabaseService from .kiaai import get_kia_model, kia_stream_and_save from ..config import settings from ..models import Chatbot, Message, BotStatus, BotType from ..messages import MESSAGES from ._utils.get_or_create_chatbot import get_or_create_chatbot from ._utils.validate_user_and_bot import validate_user_and_bot connection_kwargs = { "autocommit": False, "prepare_threshold": 0, } @tool async def web_search_tool(query: str) -> str: """Performs a web search for the given query.""" search = DuckDuckGoSearchResults() result = await search.arun(query) return result def get_model(model: str): if model.startswith("kia/"): # KIA models use KiaAIService.invoke() directly in call_model — no LangChain model needed return None if "claude" in model: # return ChatAnthropic(model=model, api_key=settings.anthropic_api_key) return ChatOpenAI( model=model, api_key=settings.liara_api_key, base_url=settings.liara_api_url ) if model in ["gpt-4o", "gpt-4o-mini","openai/gpt-4o-mini"]: return ChatOpenAI(model=model, api_key=settings.openai_api_key) # return ChatOpenAI( # model=model, # api_key=settings.liara_api_key, # base_url=settings.liara_api_url # ) if model in ["GPT-5.2", "openai/gpt-5-mini"]: return ChatOpenAI( model=model, api_key=settings.liara_api_key, base_url=settings.liara_api_url ) if model in ["Grok-3","Grok-4-Fast","x-ai/grok-3-beta","x-ai/grok-4-fast","anthropic/claude-3.5-sonnet","x-ai/grok-4","anthropic/claude-sonnet-4"]: return ChatOpenAI( model=model, api_key=settings.liara_api_key, base_url=settings.liara_api_url ) if model in ["gemini","google/gemini-3-pro-preview"]: return ChatGoogleGenerativeAI( model=model, google_api_key=settings.google_api_key ) # return ChatOpenAI( # model=model, # api_key=settings.liara_api_key, # base_url=settings.liara_api_url # ) if "deepseek" in model: return ChatOpenAI( model=model, api_key=settings.deepseek_api_key, base_url="https://api.deepseek.com", ) def document_loader(content, query: str, doc: [str | None, str | None]): loader = None if ".pdf" in doc[0]: loader = PyPDFLoader(doc[1]) elif ".docx" in doc[0]: loader = Docx2txtLoader(doc[1]) elif ".xlsx" in doc[0] or ".xls" in doc[0]: loader = UnstructuredExcelLoader(doc[1], mode="elements") elif ".txt" in doc[0]: loader = TextLoader(doc[1]) if loader: pages = [] for page in loader.lazy_load(): pages.append(page.page_content) content[0]["text"] = ( f"attached file: {' '.join(pages)}\n\n\nquery: {query}" ) content[0]["pdf_url"] = { "url": doc[0], "attachment": True, "query": query, } if os.path.exists(doc[1]): os.remove(doc[1]) return content def audio_loader(content, query: str, audio: [str | None, str | None]): parser = OpenAIWhisperParser() audio_text = "" for document in parser.lazy_parse(Blob.from_path(audio[1])): audio_text += document.page_content if query is None: content[0]["text"] = audio_text content[0]["audio_url"] = { "url": audio[0], "attachment": False, } else: content[0]["text"] = ( f"attached audio: {audio_text}\n\n\nquery: {query}" ) content[0]["audio_url"] = { "url": audio[0], "attachment": True, "query": query, } if os.path.exists(audio[1]): os.remove(audio[1]) return content def chatbot_workflow(bot, now): model = get_model(model=bot.model) trimmer = get_trimmer() if bot.web_search and model is not None: model = model.bind_tools([web_search_tool]) async def call_model(state: MessagesState): kwargs = {} if "gemini" not in bot.model: kwargs["stream_usage"] = True trimmed_messages = await trimmer.ainvoke(state["messages"]) input_messages = [ SystemMessage( f"Today: {now.isoformat()}." "You are an AI assistant designed to answer user queries conversationally, using tools when necessary." "Follow these instructions carefully: " "1. **Mathematical Formatting**: When responding to questions involving mathematics, format all mathematical expressions and formulas using Markdown with LaTeX notation for clarity and readability. Use inline LaTeX (e.g., `$x^2 + 2x + 1$`) for expressions within a sentence, and display LaTeX (e.g., `$$x^2 + 2x + 1 = 0$$`) for standalone equations or complex formulas. Ensure the formatting is compatible with Markdown renderers that support LaTeX." f"{bot.prompt if bot.web_search else ''}" "Follow these guidelines to ensure a seamless and informative conversation with the user." ) ] + trimmed_messages response = await model.ainvoke(input=input_messages, **kwargs) ai_msg_created_at = datetime.now(UTC).isoformat() response.additional_kwargs["created_at"] = ai_msg_created_at return {"messages": [response]} async def tool_node(state: MessagesState): last_message = state["messages"][-1] messages = [] for tool_call in last_message.tool_calls: if tool_call["name"] == "web_search_tool": result = await web_search_tool.ainvoke( tool_call["args"]["query"] ) messages.append( ToolMessage( content=result, tool_call_id=tool_call["id"] ) ) return {"messages": messages} workflow = StateGraph(state_schema=MessagesState) workflow.add_node("model", call_model) workflow.add_node("tool", tool_node) workflow.add_edge(START, "model") workflow.add_conditional_edges( "model", lambda state: "tool" if isinstance(state["messages"][-1], AIMessage) and state["messages"][-1].tool_calls else "end", {"tool": "tool", "end": END}, ) workflow.add_edge("tool", "model") return workflow # async def send_message( # id: int | None, # query: str, # bot_id: int, # user_id: UUID, # retry: bool, # ghost: bool, # image: str | None, # audio: [str | None, str | None], # doc: [str | None, str | None], # ): # db_service = DatabaseService(conninfo=settings.database_url, kwargs=connection_kwargs) # pool = await db_service.get_pool() # now = datetime.now(UTC) # chat_user_id = user_id # human_msg_created_at = now.isoformat() # content = [{"type": "text", "text": query}] # # --- اصلاح شده: دریافت usage_report --- # success, error, user, bot, usage_report = await validate_user_and_bot(id=id, bot_id=bot_id, user_id=user_id) # print("Usage report: ------ ", usage_report) # if not success: # yield json.dumps(error) # return # if image: # content.append({"type": "image_url", "image_url": {"url": image}}) # if audio[0]: # content = audio_loader(content, query=query, audio=audio) # if doc[0]: # content = document_loader(content, query=query, doc=doc) # title = await generate_title(content[0]["text"]) if id is None else None # try: # chatbot = await get_or_create_chatbot(id, bot_id, chat_user_id, title) # except ValueError as e: # yield str(e) # return # workflow = chatbot_workflow(bot, now) # checkpointer = AsyncPostgresSaver(pool) # try: # await checkpointer.setup() # except Exception as e: # yield json.dumps({"error": True, "detail": f"Failed to setup checkpointer: {str(e)}"}, ensure_ascii=False) # return # app = workflow.compile(checkpointer=checkpointer) # app.checkpointer = checkpointer # config = {"configurable": {"thread_id": str(chatbot.id)}} # if retry: # state = await app.aget_state(config) # messages = state.values.get("messages") # if messages: # last_human_index = max(i for i, msg in enumerate(messages) if isinstance(msg, HumanMessage)) # await app.aupdate_state( # config, # {"messages": [RemoveMessage(id=msg.id) for msg in messages[last_human_index + 1:]]}, # ) # yield json.dumps({"chat_id": chatbot.id, "chat_title": chatbot.title}, ensure_ascii=False) # input_tokens = 0 # output_tokens = 0 # human_msg = HumanMessage(content=content) # human_msg.additional_kwargs["created_at"] = human_msg_created_at # input_messages = [human_msg] # async for chunk, metadata in app.astream({"messages": input_messages}, config, stream_mode="messages"): # if isinstance(chunk, AIMessage): # if chunk.usage_metadata: # input_tokens += chunk.usage_metadata["input_tokens"] # output_tokens += chunk.usage_metadata["output_tokens"] # yield json.dumps({"content": chunk.content}, ensure_ascii=False) # await user.save() # # --- اصلاح شده: ذخیره ریز مصرف در جدول Message --- # await Message.create( # input_tokens=input_tokens, # output_tokens=output_tokens, # chatbot_id=chatbot.id, # cost=bot.cost, # cost_from_gift=usage_report["gift"], # cost_from_credit=usage_report["credit"], # cost_from_free=usage_report["free"], # ) # state = await app.aget_state(config) # messages = state.values["messages"] # human_msg_id, ai_msg_id = get_message_id(messages) # yield json.dumps({ # "credit": user.credit, # "free": user.free_credit, # "gift": user.gift_credit if user.gift_credit is not None else 0, # ارسال گیفت به فرانت # "ai_message_id": ai_msg_id, # "human_message_id": human_msg_id, # "ai_message_created_at": messages[-1].additional_kwargs.get("created_at"), # "human_message_created_at": human_msg_created_at, # }) async def send_message( id: int | None, query: str, bot_id: int, user_id: UUID, retry: bool, ghost: bool, image: str | None, audio: [str | None, str | None], doc: [str | None, str | None], ): db_service = DatabaseService(conninfo=settings.database_url, kwargs=connection_kwargs) pool = await db_service.get_pool() now = datetime.now(UTC) # تغییر ۱: همیشه از شناسه واقعی کاربر استفاده می‌کنیم (حتی در حالت روح) chat_user_id = user_id human_msg_created_at = now.isoformat() content = [{"type": "text", "text": query}] success, error, user, bot, usage_report = await validate_user_and_bot(id=id, bot_id=bot_id, user_id=user_id) print("Usage report: ------ ", usage_report) if not success: yield json.dumps(error) return if image: content.append({"type": "image_url", "image_url": {"url": image}}) if audio[0]: content = audio_loader(content, query=query, audio=audio) if doc[0]: content = document_loader(content, query=query, doc=doc) title = await generate_title(content[0]["text"]) if id is None else None try: # تغییر ۲: پاس دادن مقدار ghost به تابع سازنده چت chatbot = await get_or_create_chatbot( id=id, bot_id=bot_id, user_id=chat_user_id, title=title, is_ghost=ghost ) except ValueError as e: yield str(e) return workflow = chatbot_workflow(bot, now) checkpointer = AsyncPostgresSaver(pool) try: await checkpointer.setup() except Exception as e: yield json.dumps({"error": True, "detail": f"Failed to setup checkpointer: {str(e)}"}, ensure_ascii=False) return app = workflow.compile(checkpointer=checkpointer) app.checkpointer = checkpointer config = {"configurable": {"thread_id": str(chatbot.id)}} if retry: state = await app.aget_state(config) messages = state.values.get("messages") if messages: last_human_index = max(i for i, msg in enumerate(messages) if isinstance(msg, HumanMessage)) await app.aupdate_state( config, {"messages": [RemoveMessage(id=msg.id) for msg in messages[last_human_index + 1:]]}, ) yield json.dumps({"chat_id": chatbot.id, "chat_title": chatbot.title}, ensure_ascii=False) input_tokens = 0 output_tokens = 0 human_msg = HumanMessage(content=content) human_msg.additional_kwargs["created_at"] = human_msg_created_at input_messages = [human_msg] # ── KIA path: token-by-token streaming via kiaai.workflow ────────────── if bot.model.startswith("kia/"): try: async for item in kia_stream_and_save(app, config, human_msg, bot, now, get_trimmer()): if isinstance(item, dict): input_tokens = item["input_tokens"] output_tokens = item["output_tokens"] ai_msg_created_at = item["ai_msg_created_at"] human_msg_id, ai_msg_id = get_message_id(item["final_msgs"]) else: yield item except Exception as e: print(f"[KiaAI] Streaming error: {e}") yield json.dumps({"error": True, "detail": str(e)}, ensure_ascii=False) return await user.save() await Message.create( input_tokens=input_tokens, output_tokens=output_tokens, chatbot_id=chatbot.id, cost=bot.cost, cost_from_gift=usage_report["gift"], cost_from_credit=usage_report["credit"], cost_from_free=usage_report["free"], ) yield json.dumps({ "credit": user.credit, "free": user.free_credit, "gift": user.gift_credit if user.gift_credit is not None else 0, "ai_message_id": ai_msg_id, "human_message_id": human_msg_id, "ai_message_created_at": ai_msg_created_at, "human_message_created_at": human_msg_created_at, }) return # ── Non-KIA path (LangGraph astream) ──────────────────────────────────── async for chunk, metadata in app.astream({"messages": input_messages}, config, stream_mode="messages"): try: if isinstance(chunk, AIMessage): if chunk.usage_metadata: input_tokens += chunk.usage_metadata["input_tokens"] output_tokens += chunk.usage_metadata["output_tokens"] yield json.dumps({"content": chunk.content}, ensure_ascii=False) except Exception as e: print(f"[send_message] Error processing chunk: {e}") yield json.dumps({"error": True, "detail": str(e)}, ensure_ascii=False) return await user.save() # تغییر ۳: ذخیره payer_id برای گزارش‌گیری دقیق مالی await Message.create( input_tokens=input_tokens, output_tokens=output_tokens, chatbot_id=chatbot.id, cost=bot.cost, cost_from_gift=usage_report["gift"], cost_from_credit=usage_report["credit"], cost_from_free=usage_report["free"], ) state = await app.aget_state(config) messages = state.values["messages"] human_msg_id, ai_msg_id = get_message_id(messages) yield json.dumps({ "credit": user.credit, "free": user.free_credit, "gift": user.gift_credit if user.gift_credit is not None else 0, "ai_message_id": ai_msg_id, "human_message_id": human_msg_id, "ai_message_created_at": messages[-1].additional_kwargs.get("created_at"), "human_message_created_at": human_msg_created_at, }) async def get_chats( user_id: UUID, query: str | None, archive: bool, date: date | None, page: int, type: BotType | None = None, ): limit = 20 offset = (page - 1) * limit filters = Q(archive=archive, user_id=user_id, bot__status=BotStatus.CONFIRMED, is_deleted=False,is_ghost=False) if query is not None: filters &= Q(title__icontains=query) if date is not None: start_date = datetime.combine(date, datetime.min.time()) start_date -= timedelta(hours=3, minutes=30) filters &= Q(created_at__gte=start_date, created_at__lte=start_date + timedelta(days=1)) if type is not None: filters &= Q(bot__type=type) total_count = await Chatbot.filter(filters).count() chats = ( await Chatbot.filter(filters) .prefetch_related("bot", "bot__category") .order_by("-created_at") .offset(offset) .limit(limit) ) return { "chats": chats, "page": page, "total_count": total_count, "last_page": math.ceil(total_count / limit), } async def get_chat(id: int, user_id: UUID): print(f"=== get_chat START | id: {id}, user_id: {user_id} ===") db_service = DatabaseService(conninfo=settings.database_url, kwargs=connection_kwargs) pool = await db_service.get_pool() print("[1] DB Pool initialized") try: print(f"[2] Fetching Chatbot with id={id} and user_id={user_id}...") chatbot: Chatbot = await Chatbot.get_or_none( id=id, user_id=user_id, bot__status=BotStatus.CONFIRMED, is_deleted=False ).prefetch_related("bot", "bot__category") if chatbot is None: print(f"❌ [Error] Chatbot not found for id={id}") raise HTTPException(404, detail=MESSAGES["chat"]["not_found"]) print(f"✅ [3] Chatbot found: {chatbot.id}") checkpointer = AsyncPostgresSaver(pool) await checkpointer.setup() print("[4] Checkpointer setup done") config = {"configurable": {"thread_id": str(id)}} print(f"[5] Fetching history with config: {config}") history = await checkpointer.aget_tuple(config) print(f"[6] History retrieved (checkpoint found: {history is not None})") if history: # چاپ جزئیات بیشتر در صورت نیاز (مثلاً آخرین پیام) print(f" History metadata: {history.metadata if hasattr(history, 'metadata') else 'N/A'}") print("=== get_chat END (Success) ===") return history, chatbot except HTTPException as e: print(f"⚠️ [Handled Exception] HTTP {e.status_code}: {e.detail}") raise e except Exception as e: print(f"❌ [Critical Error] in get_chat: {str(e)}") raise HTTPException(status_code=500, detail=str(e)) async def delete_chat(id: int, user_id: UUID): db_service = DatabaseService(conninfo=settings.database_url, kwargs=connection_kwargs) pool = await db_service.get_pool() try: chatbot: Chatbot = await Chatbot.get_or_none(id=id, user_id=user_id) if chatbot is None: raise HTTPException(404, detail=MESSAGES["chat"]["not_found"]) checkpointer = AsyncPostgresSaver(pool) await checkpointer.setup() config = {"configurable": {"thread_id": str(id)}} workflow = StateGraph(state_schema=MessagesState) app = workflow.compile(checkpointer=checkpointer) state = await app.aget_state(config) messages = state.values["messages"] await app.aupdate_state( config, {"messages": [RemoveMessage(id=message.id) for message in messages]}, ) # await chatbot.delete() chatbot.is_deleted = True await chatbot.save() await Message.filter(chatbot_id=chatbot.id).update(is_deleted=True) except HTTPException as e: raise e except Exception as e: raise HTTPException(status_code=500, detail=str(e)) async def delete_chats(user_id: UUID, archive: bool): db_service = DatabaseService(conninfo=settings.database_url, kwargs=connection_kwargs) pool = await db_service.get_pool() try: chatbots: list[Chatbot] = await Chatbot.filter(user_id=user_id, archive=archive).all() checkpointer = AsyncPostgresSaver(pool) await checkpointer.setup() for chatbot in chatbots: config = {"configurable": {"thread_id": str(chatbot.id)}} workflow = StateGraph(state_schema=MessagesState) app = workflow.compile(checkpointer=checkpointer) state = await app.aget_state(config) messages = state.values.get("messages", None) if messages: await app.aupdate_state( config, {"messages": [RemoveMessage(id=message.id) for message in messages]}, ) # await chatbot.delete() chatbot.is_deleted = True await chatbot.save() await Message.filter(chatbot_id=chatbot.id).update(is_deleted=True) except HTTPException as e: raise e except Exception as e: raise HTTPException(status_code=500, detail=str(e)) async def delete_message(id: int, message_id: str, user_id: UUID): db_service = DatabaseService(conninfo=settings.database_url, kwargs=connection_kwargs) pool = await db_service.get_pool() try: chatbot: Chatbot = await Chatbot.get_or_none(id=id, user_id=user_id) if chatbot is None: raise HTTPException(404, detail=MESSAGES["chat"]["not_found"]) checkpointer = AsyncPostgresSaver(pool) await checkpointer.setup() config = {"configurable": {"thread_id": str(id)}} workflow = StateGraph(state_schema=MessagesState) app = workflow.compile(checkpointer=checkpointer) await app.aupdate_state(config, {"messages": RemoveMessage(id=message_id)}) except HTTPException as e: raise e except ValueError as e: raise HTTPException(status_code=404, detail=str(e)) except Exception as e: raise HTTPException(status_code=500, detail=str(e)) async def like_message(id: int, message_id: str, user_id: UUID, like: bool | None): db_service = DatabaseService(conninfo=settings.database_url, kwargs=connection_kwargs) pool = await db_service.get_pool() try: chatbot: Chatbot = await Chatbot.get_or_none(id=id, user_id=user_id) if chatbot is None: raise HTTPException(404, detail="Chat not found") checkpointer = AsyncPostgresSaver(pool) await checkpointer.setup() config = {"configurable": {"thread_id": str(id)}} workflow = StateGraph(state_schema=MessagesState) app = workflow.compile(checkpointer=checkpointer) state = await app.aget_state(config) messages = state.values["messages"] message = next((message for message in messages if message.id == message_id), None) if message is None: raise HTTPException(404, detail=MESSAGES["message"]["not_found"]) message.additional_kwargs["like"] = like await app.aupdate_state(config, {"messages": message}) except HTTPException as e: raise e except ValueError as e: raise HTTPException(status_code=404, detail=str(e)) except Exception as e: raise HTTPException(status_code=500, detail=str(e)) async def edit_title(id: int, user_id: UUID, title: str): chatbot: Chatbot = await Chatbot.get_or_none(id=id, user_id=user_id) if chatbot is None: raise HTTPException(404, detail=MESSAGES["chat"]["not_found"]) chatbot.title = title await chatbot.save() async def archive(id: int, user_id: UUID, archive: bool): chatbot: Chatbot = await Chatbot.get_or_none(id=id, user_id=user_id) if chatbot is None: raise HTTPException(404, detail=MESSAGES["chat"]["not_found"]) chatbot.archive = archive await chatbot.save() async def related_questions(query: str): model = ChatOpenAI(model="gpt-4o-mini", api_key=settings.openai_api_key) instruction = "Generate 3 related questions based on the following query." res = await model.ainvoke([SystemMessage(content=instruction), HumanMessage(content=query)]) related_questions = [q.strip()[3:] for q in res.content.split("\n") if q.strip()] return related_questions async def generate_title(query: str): model = ChatOpenAI(model="gpt-4o-mini", api_key=settings.openai_api_key) instruction = "Generate a concise title in Farsi from the following prompt." res = await model.ainvoke([SystemMessage(content=instruction), HumanMessage(content=query)]) return res.content def get_trimmer(): token_counter = get_model("gpt-4o-mini") return trim_messages( max_tokens=6000, strategy="last", token_counter=token_counter, include_system=True, allow_partial=True, start_on="human", ) def get_message_id(messages): human_messages = [msg for msg in messages if isinstance(msg, HumanMessage)] ai_messages = [msg for msg in messages if isinstance(msg, AIMessage) and not msg.tool_calls] return human_messages[-1].id, ai_messages[-1].id