diff --git a/backend/src/api/chat.py b/backend/src/api/chat.py index f56b683..9dcd5a1 100644 --- a/backend/src/api/chat.py +++ b/backend/src/api/chat.py @@ -1,11 +1,13 @@ """ 聊天对话API """ +import logging + from fastapi import APIRouter, Depends, HTTPException, status from fastapi.responses import StreamingResponse from sqlalchemy.orm import Session -from pydantic import BaseModel -from typing import List, Optional, Dict, Any +from pydantic import BaseModel, Field +from typing import List, Optional, Dict, Any, Literal import json import asyncio @@ -17,6 +19,8 @@ from ..rag.chains import create_rag_chain from ..rag.conversation_chains import create_conversation_chain from ..llm.siliconflow import get_llm_client +logger = logging.getLogger(__name__) + router = APIRouter(prefix="/chat", tags=["聊天"]) @@ -33,11 +37,11 @@ def get_user_id_by_username(db: Session, username: str) -> int: class ChatRequest(BaseModel): """聊天请求模型""" - message: str + message: str = Field(..., min_length=1, max_length=5000, description="用户消息") session_id: Optional[int] = None - mode: str = "normal" # normal, rag - knowledge_base_ids: Optional[List[int]] = None - model: Optional[str] = None # 模型ID,如 deepseek-ai/DeepSeek-V3, Qwen/QwQ-32B + mode: Literal["normal", "rag"] = "normal" + knowledge_base_ids: Optional[List[int]] = Field(None, description="RAG模式使用的知识库ID列表") + model: Optional[str] = None class ChatResponse(BaseModel): @@ -105,24 +109,24 @@ async def send_message( new_title = request.message.strip()[:30] + "..." if len(request.message.strip()) > 30 else request.message.strip() session.title = new_title db.commit() - print(f"[DEBUG-AUTO-TITLE] 自动更新会话标题: {new_title}") - + logger.debug(f"自动更新会话标题: {new_title}") + # 根据模式运行不同的问答工作流 if request.mode == "rag": - print(f"[DEBUG-RAG] 非流式RAG模式") - print(f"[DEBUG-RAG] 接收到的knowledge_base_ids: {request.knowledge_base_ids}") - print(f"[DEBUG-RAG] knowledge_base_ids类型: {type(request.knowledge_base_ids)}") + logger.debug(f"非流式RAG模式") + logger.debug(f"接收到的knowledge_base_ids: {request.knowledge_base_ids}") + logger.debug(f"knowledge_base_ids类型: {type(request.knowledge_base_ids)}") if request.knowledge_base_ids: - print(f"[DEBUG-RAG] 第一个ID: {request.knowledge_base_ids[0]}, 类型: {type(request.knowledge_base_ids[0])}") - + logger.debug(f"第一个ID: {request.knowledge_base_ids[0]}, 类型: {type(request.knowledge_base_ids[0])}") + # 使用LangChain 1.0 RAG链 - print(f"[DEBUG-RAG] 模型: {request.model}") + logger.debug(f"模型: {request.model}") rag_chain = create_rag_chain(knowledge_base_ids=request.knowledge_base_ids, model=request.model) result = rag_chain.invoke(request.message) else: # 普通模式:使用LangChain 1.0对话链 - print(f"[DEBUG-CHAT] 普通模式 - 使用LangChain对话链") - print(f"[DEBUG-CHAT] 模型: {request.model}") + logger.debug(f"普通模式 - 使用LangChain对话链") + logger.debug(f"模型: {request.model}") conversation_chain = create_conversation_chain(model=request.model) # 获取聊天历史 @@ -175,7 +179,7 @@ async def stream_message( db = SessionLocal() try: # 打印请求参数调试信息 - print(f"[DEBUG-CHAT] 接收流式请求 - mode: {request.mode}, knowledge_base_ids: {request.knowledge_base_ids}, model: {request.model}, message: {request.message[:50]}...") + logger.debug(f"接收流式请求 - mode: {request.mode}, knowledge_base_ids: {request.knowledge_base_ids}, model: {request.model}, message: {request.message[:50]}...") # 获取或创建会话 if request.session_id: @@ -215,7 +219,7 @@ async def stream_message( new_title = request.message.strip()[:30] + "..." if len(request.message.strip()) > 30 else request.message.strip() session.title = new_title db.commit() - print(f"[DEBUG-AUTO-TITLE] 自动更新会话标题: {new_title}") + logger.debug(f"自动更新会话标题: {new_title}") # 获取聊天历史 chat_history = [] @@ -231,12 +235,12 @@ async def stream_message( # 根据模式选择不同的处理方式 if request.mode == "rag": - print(f"[DEBUG-RAG] 使用LangChain 1.0 RAG链") - print(f"[DEBUG-RAG] 接收到的knowledge_base_ids: {request.knowledge_base_ids}") - print(f"[DEBUG-RAG] knowledge_base_ids类型: {type(request.knowledge_base_ids)}") - print(f"[DEBUG-RAG] 模型: {request.model}") + logger.debug(f"使用LangChain 1.0 RAG链") + logger.debug(f"接收到的knowledge_base_ids: {request.knowledge_base_ids}") + logger.debug(f"knowledge_base_ids类型: {type(request.knowledge_base_ids)}") + logger.debug(f"模型: {request.model}") if request.knowledge_base_ids: - print(f"[DEBUG-RAG] 第一个ID: {request.knowledge_base_ids[0]}, 类型: {type(request.knowledge_base_ids[0])}") + logger.debug(f"第一个ID: {request.knowledge_base_ids[0]}, 类型: {type(request.knowledge_base_ids[0])}") # 创建RAG链 rag_chain = create_rag_chain( @@ -288,8 +292,8 @@ async def stream_message( else: # 普通模式:使用LangChain 1.0对话链 - print(f"[DEBUG-CHAT] 普通模式 - 使用LangChain流式对话链") - print(f"[DEBUG-CHAT] 模型: {request.model}") + logger.debug(f"普通模式 - 使用LangChain流式对话链") + logger.debug(f"模型: {request.model}") conversation_chain = create_conversation_chain(model=request.model) # 使用带思考过程的流式输出 @@ -739,7 +743,7 @@ async def run_rag_workflow_with_context(question: str, session_id: int, db: Sess return result except Exception as e: - print(f"RAG工作流执行失败: {str(e)}") + logger.error(f"RAG工作流执行失败: {str(e)}") # 降级到基础问答 rag_chain = create_rag_chain(model=model) return rag_chain.invoke(question) diff --git a/backend/src/graph/nodes.py b/backend/src/graph/nodes.py index 0d03c26..6ef1d1b 100644 --- a/backend/src/graph/nodes.py +++ b/backend/src/graph/nodes.py @@ -2,6 +2,7 @@ LangGraph节点定义 """ import json +import logging import re from typing import Dict, Any, List, TypedDict from langchain.schema import HumanMessage, AIMessage, SystemMessage @@ -9,6 +10,8 @@ from ..llm.siliconflow import get_llm_client from ..rag.retrievers import KnowledgeBaseRetriever from ..rag.vector_store import get_vector_store +logger = logging.getLogger(__name__) + def _parse_json_from_response(text: str) -> dict: """从LLM响应中提取JSON,兼容markdown代码块包裹的情况""" @@ -73,7 +76,7 @@ def analyze_question_node(state: GraphState) -> GraphState: return new_state except Exception as e: - print(f"问题分析失败: {str(e)}") + logger.error(f"问题分析失败: {str(e)}") new_state = state.copy() new_state["metadata"] = {**state["metadata"], "analysis": { "question_type": "通用", @@ -138,9 +141,7 @@ def retrieve_knowledge_node(state: GraphState) -> GraphState: return new_state except Exception as e: - print(f"知识检索失败: {str(e)}") - import traceback - print(traceback.format_exc()) + logger.error(f"知识检索失败: {str(e)}", exc_info=True) new_state = state.copy() new_state["retrieved_docs"] = [] new_state["sources"] = [] @@ -186,7 +187,7 @@ def generate_answer_node(state: GraphState) -> GraphState: return new_state except Exception as e: - print(f"答案生成失败: {str(e)}") + logger.error(f"答案生成失败: {str(e)}") new_state = state.copy() new_state["answer"] = "抱歉,我无法生成合适的回答。请稍后重试。" return new_state @@ -208,5 +209,5 @@ def format_response_node(state: GraphState) -> GraphState: return new_state except Exception as e: - print(f"响应格式化失败: {str(e)}") + logger.error(f"响应格式化失败: {str(e)}") return state diff --git a/backend/src/graph/qa_graph.py b/backend/src/graph/qa_graph.py index 7a8d285..bc6f4c1 100644 --- a/backend/src/graph/qa_graph.py +++ b/backend/src/graph/qa_graph.py @@ -1,10 +1,13 @@ """ 问答LangGraph工作流 """ +import logging from typing import Dict, Any, Optional, List from langgraph.graph import StateGraph, END from .nodes import GraphState, analyze_question_node, retrieve_knowledge_node, generate_answer_node, format_response_node +logger = logging.getLogger(__name__) + def create_qa_graph() -> StateGraph: """创建问答图""" @@ -57,7 +60,7 @@ def run_qa_workflow(question: str, knowledge_base_ids: Optional[List[int]] = Non } except Exception as e: - print(f"问答工作流执行失败: {str(e)}") + logger.error(f"问答工作流执行失败: {str(e)}") return { "answer": "抱歉,处理您的问题时出现了错误。请稍后重试。", "sources": [], diff --git a/backend/src/llm/siliconflow.py b/backend/src/llm/siliconflow.py index 2a7dd41..963d579 100644 --- a/backend/src/llm/siliconflow.py +++ b/backend/src/llm/siliconflow.py @@ -1,6 +1,7 @@ """ 大模型API集成 — 支持 SiliconFlow 和 DeepSeek 官方 + 视觉模型 """ +import logging import os import base64 from typing import List, Dict, Any, Optional, AsyncGenerator, Tuple @@ -12,6 +13,8 @@ import openai from ..core.config import get_settings +logger = logging.getLogger(__name__) + settings = get_settings() # 图片描述提示词 @@ -56,7 +59,7 @@ class SiliconFlowLLM: api_key, base_url, resolved_model = _resolve_provider(raw_model) self.model_name = resolved_model - print(f"[LLM] 模型: {resolved_model}, API: {base_url}") + logger.info(f"模型: {resolved_model}, API: {base_url}") self.llm = ChatOpenAI( model=resolved_model, @@ -210,7 +213,7 @@ class SiliconFlowLLM: return response.choices[0].message.content or "" except Exception as e: - print(f"[VLM] 描述失败 attempt={attempt+1}: {e}") + logger.warning(f"描述失败 attempt={attempt+1}: {e}") if attempt < max_retries - 1: await asyncio.sleep(2 ** attempt) diff --git a/backend/src/migrations/migrate_orphaned_documents.py b/backend/src/migrations/migrate_orphaned_documents.py index 9c83d52..56c33b2 100644 --- a/backend/src/migrations/migrate_orphaned_documents.py +++ b/backend/src/migrations/migrate_orphaned_documents.py @@ -1,12 +1,15 @@ """ 迁移孤立文档到默认知识库 """ +import logging from sqlalchemy.orm import Session from ..core.database import get_db from ..models.user import User from ..models.knowledge_base import KnowledgeBase from ..models.document import Document +logger = logging.getLogger(__name__) + def migrate_orphaned_documents(): """将没有知识库的文档迁移到用户的默认知识库""" @@ -40,12 +43,12 @@ def migrate_orphaned_documents(): doc.knowledge_base_id = default_kb.id db.commit() - print(f"用户 {user.username} 的 {len(orphaned_docs)} 个文档已迁移到默认知识库") + logger.info(f"用户 {user.username} 的 {len(orphaned_docs)} 个文档已迁移到默认知识库") - print("孤立文档迁移完成") + logger.info("孤立文档迁移完成") except Exception as e: - print(f"迁移失败: {str(e)}") + logger.error(f"迁移失败: {str(e)}") db.rollback() finally: db.close() diff --git a/backend/src/rag/chains.py b/backend/src/rag/chains.py index c7c01be..2fc83b0 100644 --- a/backend/src/rag/chains.py +++ b/backend/src/rag/chains.py @@ -1,6 +1,7 @@ """ RAG检索链(LangChain 1.0) """ +import logging from typing import List, Optional, Dict, Any from langchain_core.runnables import RunnablePassthrough, RunnableLambda from langchain_core.prompts import ChatPromptTemplate @@ -11,6 +12,8 @@ from .retrievers import KnowledgeBaseRetriever from .vector_store import get_vector_store from ..llm.siliconflow import get_llm_client +logger = logging.getLogger(__name__) + class RAGChain: """RAG问答链(LangChain 1.0标准API)""" @@ -31,7 +34,7 @@ class RAGChain: score_threshold: 相似度阈值 model: 可选的模型名称,如 deepseek-ai/DeepSeek-V3, Qwen/QwQ-32B """ - print(f"[DEBUG-RAGChain] 初始化,knowledge_base_ids: {knowledge_base_ids}, model: {model}") + logger.debug(f"初始化RAG链,knowledge_base_ids: {knowledge_base_ids}, model: {model}") _sf_client = get_llm_client(model=model) self.llm = _sf_client.llm self.client = _sf_client @@ -46,7 +49,7 @@ class RAGChain: search_kwargs={"k": k}, score_threshold=score_threshold ) - print(f"[DEBUG-RAGChain] 检索器创建完成,knowledge_base_ids: {self.retriever.knowledge_base_ids}") + logger.debug(f"检索器创建完成,knowledge_base_ids: {self.retriever.knowledge_base_ids}") # 创建Prompt self.prompt = create_rag_prompt() diff --git a/backend/src/rag/document_loaders.py b/backend/src/rag/document_loaders.py index 7490b67..d0becf6 100644 --- a/backend/src/rag/document_loaders.py +++ b/backend/src/rag/document_loaders.py @@ -1,6 +1,7 @@ """ LangChain 1.0 文档加载器封装 + PDF图片提取 """ +import logging from typing import List, Optional, Dict from pathlib import Path import fitz # pymupdf @@ -12,6 +13,8 @@ from langchain_community.document_loaders import ( ) from langchain_core.documents import Document +logger = logging.getLogger(__name__) + class PDFImageExtractor: """使用pymupdf从PDF中提取内嵌图片""" @@ -59,7 +62,7 @@ class PDFImageExtractor: "size": len(image_bytes), }) except Exception as e: - print(f"[ImageExtractor] 提取图片失败 page={page_num+1} img={img_idx}: {e}") + logger.warning(f"提取图片失败 page={page_num+1} img={img_idx}: {e}") continue doc.close() diff --git a/backend/src/rag/embeddings.py b/backend/src/rag/embeddings.py index c234a0f..873f20d 100644 --- a/backend/src/rag/embeddings.py +++ b/backend/src/rag/embeddings.py @@ -1,12 +1,15 @@ """ 嵌入模型管理 """ +import logging from typing import List from sentence_transformers import SentenceTransformer import numpy as np from ..core.config import get_settings +logger = logging.getLogger(__name__) + settings = get_settings() @@ -23,7 +26,7 @@ class EmbeddingModel: """加载嵌入模型""" try: self.model = SentenceTransformer(self.model_name) - print(f"嵌入模型 {self.model_name} 加载成功") + logger.info(f"嵌入模型 {self.model_name} 加载成功") except Exception as e: raise Exception(f"嵌入模型加载失败: {str(e)}") diff --git a/backend/src/rag/retriever.py b/backend/src/rag/retriever.py index 4b175e1..1c6bfcd 100644 --- a/backend/src/rag/retriever.py +++ b/backend/src/rag/retriever.py @@ -1,10 +1,13 @@ """ 文档检索器 """ +import logging from typing import List, Dict, Any, Optional from .vector_store import get_vector_store from .embeddings import get_embedding_model +logger = logging.getLogger(__name__) + class DocumentRetriever: """文档检索器""" @@ -23,7 +26,7 @@ class DocumentRetriever: ) -> List[Dict[str, Any]]: """检索相关文档""" try: - print(f"[DEBUG-RETRIEVER] 开始检索 - query: {query[:50]}..., filter_metadata: {filter_metadata}") + logger.debug(f"开始检索 - query: {query[:50]}..., filter_metadata: {filter_metadata}") # 执行向量搜索 results = self.vector_store.search( @@ -31,19 +34,19 @@ class DocumentRetriever: n_results=top_k, filter_metadata=filter_metadata ) - print(f"[DEBUG-RETRIEVER] 向量搜索返回结果数量: {len(results)}") + logger.debug(f"向量搜索返回结果数量: {len(results)}") # 过滤低分结果 filtered_results = [ result for result in results if result.get("distance", 1.0) <= (1 - score_threshold) ] - print(f"[DEBUG-RETRIEVER] 过滤后结果数量: {len(filtered_results)}") + logger.debug(f"过滤后结果数量: {len(filtered_results)}") return filtered_results except Exception as e: - print(f"文档检索失败: {str(e)}") + logger.error(f"文档检索失败: {str(e)}") return [] def retrieve_by_document_id(self, document_id: str) -> List[Dict[str, Any]]: @@ -56,7 +59,7 @@ class DocumentRetriever: ) return results except Exception as e: - print(f"按文档ID检索失败: {str(e)}") + logger.error(f"按文档ID检索失败: {str(e)}") return [] def get_relevant_context(self, query: str, max_length: int = 2000) -> str: diff --git a/backend/src/rag/retrievers.py b/backend/src/rag/retrievers.py index 7d0b3ba..b8653b2 100644 --- a/backend/src/rag/retrievers.py +++ b/backend/src/rag/retrievers.py @@ -1,10 +1,13 @@ """ 自定义知识库检索器(LangChain 1.0) """ +import logging from typing import List, Optional -import math + +logger = logging.getLogger(__name__) from langchain_core.documents import Document from langchain_core.callbacks.manager import CallbackManagerForRetrieverRun +from .score_utils import convert_distance_to_score from langchain_core.retrievers import BaseRetriever class KnowledgeBaseRetriever(BaseRetriever): @@ -23,16 +26,16 @@ class KnowledgeBaseRetriever(BaseRetriever): run_manager: CallbackManagerForRetrieverRun ) -> List[Document]: """获取相关文档(LangChain 1.0标准接口)""" - print(f"[DEBUG-Retriever] 查询: {query}") - print(f"[DEBUG-Retriever] knowledge_base_ids: {self.knowledge_base_ids}") + logger.debug(f"查询: {query}") + logger.debug(f"knowledge_base_ids: {self.knowledge_base_ids}") # 构建知识库过滤条件 filter_dict = None if self.knowledge_base_ids: filter_dict = {"knowledge_base_id": {"$in": self.knowledge_base_ids}} - print(f"[DEBUG-Retriever] 构建的过滤条件: {filter_dict}") + logger.debug(f"构建的过滤条件: {filter_dict}") else: - print(f"[DEBUG-Retriever] 没有知识库ID,不进行过滤") + logger.debug("没有知识库ID,不进行过滤") # 执行搜索 if self.search_type == "similarity": @@ -41,21 +44,21 @@ class KnowledgeBaseRetriever(BaseRetriever): k=self.search_kwargs.get("k", 5), filter=filter_dict ) - print(f"[DEBUG-Retriever] 搜索返回文档数量: {len(docs_and_scores)}") + logger.debug(f"搜索返回文档数量: {len(docs_and_scores)}") # 打印每个文档的知识库ID for i, (doc, distance) in enumerate(docs_and_scores): kb_id = doc.metadata.get("knowledge_base_id", "未知") - print(f"[DEBUG-Retriever] 文档{i+1}: knowledge_base_id={kb_id}, distance={distance:.4f}") + logger.debug(f"文档{i+1}: knowledge_base_id={kb_id}, distance={distance:.4f}") # 转换距离为分数并过滤 filtered_docs = [] for doc, distance in docs_and_scores: - score = self._convert_distance_to_score(distance) + score = convert_distance_to_score(distance) if score > self.score_threshold: filtered_docs.append(doc) - print(f"[DEBUG-Retriever] 过滤后文档数量: {len(filtered_docs)}") + logger.debug(f"过滤后文档数量: {len(filtered_docs)}") return filtered_docs elif self.search_type == "mmr": @@ -67,16 +70,3 @@ class KnowledgeBaseRetriever(BaseRetriever): ) return [] - - def _convert_distance_to_score(self, distance: float) -> float: - """将ChromaDB的distance转换为0-1范围的相似度分数""" - # 内积距离(负值) - if distance < 0: - return (1 + distance) / 2 - - # 大距离使用对数缩放 - if distance > 100: - return 1 / (1 + math.log(distance)) - - # 标准距离转换 - return 1 / (1 + distance) diff --git a/backend/src/rag/vector_store.py b/backend/src/rag/vector_store.py index 0f46da5..fceec2e 100644 --- a/backend/src/rag/vector_store.py +++ b/backend/src/rag/vector_store.py @@ -2,6 +2,7 @@ LangChain 1.0 向量存储封装 """ import os +import logging import json import hashlib from typing import List, Dict, Any, Optional @@ -12,6 +13,8 @@ from langchain_core.documents import Document as LangChainDocument from ..core.config import get_settings from .embeddings import get_embedding_model +logger = logging.getLogger(__name__) + settings = get_settings() @@ -40,7 +43,7 @@ class VectorStore: self.vectorstore.add_documents(documents) return True except Exception as e: - print(f"添加文档失败: {str(e)}") + logger.error(f"添加文档失败: {str(e)}") return False def as_retriever(self, **kwargs): @@ -54,7 +57,7 @@ class VectorStore: filter: Optional[Dict] = None ): """相似度搜索(带分数)""" - print(f"[DEBUG-VectorStore] 查询参数 - k: {k}, filter: {filter}") + logger.debug(f"查询参数 - k: {k}, filter: {filter}") result = self.vectorstore.similarity_search_with_score( query=query, @@ -62,7 +65,7 @@ class VectorStore: filter=filter ) - print(f"[DEBUG-VectorStore] 返回结果数量: {len(result)}") + logger.debug(f"返回结果数量: {len(result)}") return result def delete_by_document_id(self, document_id: int) -> bool: @@ -73,7 +76,7 @@ class VectorStore: ) return True except Exception as e: - print(f"删除向量数据失败: {str(e)}") + logger.error(f"删除向量数据失败: {str(e)}") return False def max_marginal_relevance_search( diff --git a/backend/src/services/analytics_service.py b/backend/src/services/analytics_service.py index 8492210..7949ae3 100644 --- a/backend/src/services/analytics_service.py +++ b/backend/src/services/analytics_service.py @@ -2,12 +2,14 @@ 学习分析服务 提供用户学习数据的统计和分析功能 """ - +import logging from typing import Dict, Any, List from sqlalchemy.orm import Session from sqlalchemy import func, distinct, and_ from datetime import datetime, timedelta +logger = logging.getLogger(__name__) + from ..models.user import User from ..models.chat import ChatSession, ChatMessage from ..models.document import Document @@ -65,7 +67,7 @@ class AnalyticsService: } except Exception as e: - print(f"[ERROR] 获取用户统计数据失败: {str(e)}") + logger.error(f"获取用户统计数据失败: {str(e)}") raise Exception(f"获取统计数据失败: {str(e)}") def get_user_learning_trends(self, user_id: int, days: int = 30) -> List[Dict[str, Any]]: @@ -107,7 +109,7 @@ class AnalyticsService: return trends except Exception as e: - print(f"[ERROR] 获取学习趋势数据失败: {str(e)}") + logger.error(f"获取学习趋势数据失败: {str(e)}") raise Exception(f"获取学习趋势失败: {str(e)}") def get_popular_questions(self, user_id: int, limit: int = 10) -> List[Dict[str, Any]]: @@ -141,7 +143,7 @@ class AnalyticsService: return popular_questions except Exception as e: - print(f"[ERROR] 获取热门问题失败: {str(e)}") + logger.error(f"获取热门问题失败: {str(e)}") raise Exception(f"获取热门问题失败: {str(e)}") def get_knowledge_coverage(self, user_id: int) -> List[Dict[str, Any]]: @@ -160,7 +162,7 @@ class AnalyticsService: return coverage_data except Exception as e: - print(f"[ERROR] 获取知识覆盖度失败: {str(e)}") + logger.error(f"获取知识覆盖度失败: {str(e)}") raise Exception(f"获取知识覆盖度失败: {str(e)}") def get_learning_report(self, user_id: int) -> Dict[str, Any]: @@ -211,5 +213,5 @@ class AnalyticsService: } except Exception as e: - print(f"[ERROR] 获取学习报告失败: {str(e)}") + logger.error(f"获取学习报告失败: {str(e)}") raise Exception(f"获取学习报告失败: {str(e)}") diff --git a/backend/src/services/auth_service.py b/backend/src/services/auth_service.py index 688d74f..96eac22 100644 --- a/backend/src/services/auth_service.py +++ b/backend/src/services/auth_service.py @@ -1,6 +1,7 @@ """ 用户认证服务 """ +import logging from datetime import datetime, timedelta from typing import Optional, Dict, Any from sqlalchemy.orm import Session @@ -10,6 +11,8 @@ from ..models.user import User from ..core.security import get_password_hash, verify_password, create_access_token from ..core.config import get_settings +logger = logging.getLogger(__name__) + settings = get_settings() @@ -58,7 +61,7 @@ class AuthService: except Exception as e: self.db.rollback() - print(f"创建用户失败: {str(e)}") + logger.error(f"创建用户失败: {str(e)}") return None def authenticate_user(self, username: str, password: str) -> Optional[User]: @@ -85,7 +88,7 @@ class AuthService: return user except Exception as e: - print(f"用户认证失败: {str(e)}") + logger.error(f"用户认证失败: {str(e)}") return None def get_user_by_username(self, username: str) -> Optional[User]: @@ -117,7 +120,7 @@ class AuthService: except Exception as e: self.db.rollback() - print(f"更新用户失败: {str(e)}") + logger.error(f"更新用户失败: {str(e)}") return None def deactivate_user(self, user_id: int) -> bool: @@ -133,7 +136,7 @@ class AuthService: except Exception as e: self.db.rollback() - print(f"停用用户失败: {str(e)}") + logger.error(f"停用用户失败: {str(e)}") return False def change_password(self, user_id: int, old_password: str, new_password: str) -> bool: @@ -152,7 +155,7 @@ class AuthService: except Exception as e: self.db.rollback() - print(f"修改密码失败: {str(e)}") + logger.error(f"修改密码失败: {str(e)}") return False def create_access_token_for_user(self, user: User) -> Dict[str, Any]: @@ -193,7 +196,7 @@ class AuthService: } except Exception as e: - print(f"获取用户统计失败: {str(e)}") + logger.error(f"获取用户统计失败: {str(e)}") return { "total_users": 0, "active_users": 0, diff --git a/backend/src/services/book_content_service.py b/backend/src/services/book_content_service.py index 4374c98..082d5d8 100644 --- a/backend/src/services/book_content_service.py +++ b/backend/src/services/book_content_service.py @@ -3,6 +3,7 @@ 根据数据库中的行号范围从LaTeX文件动态读取内容 层级结构:书籍 -> Chapter -> Section -> Subsection(知识点) """ +import logging from pathlib import Path from typing import Optional from sqlalchemy.orm import Session @@ -10,6 +11,8 @@ from sqlalchemy.orm import Session from ..models.book_structure import Chapter, Section, Subsection from .latex_parser import LaTeXParser +logger = logging.getLogger(__name__) + class BookContentService: """书籍内容服务""" @@ -59,7 +62,7 @@ class BookContentService: ) return content except Exception as e: - print(f"读取章节内容失败: {e}") + logger.error(f"读取章节内容失败: {e}") return None def get_section_content(self, db: Session, section_id: int) -> Optional[str]: @@ -91,7 +94,7 @@ class BookContentService: ) return content except Exception as e: - print(f"读取节内容失败: {e}") + logger.error(f"读取节内容失败: {e}") return None def get_subsection_content(self, db: Session, subsection_id: int) -> Optional[str]: @@ -127,6 +130,6 @@ class BookContentService: ) return content except Exception as e: - print(f"读取小节内容失败: {e}") + logger.error(f"读取小节内容失败: {e}") return None diff --git a/backend/src/services/file_watcher_service.py b/backend/src/services/file_watcher_service.py index 8f5c300..40c3f63 100644 --- a/backend/src/services/file_watcher_service.py +++ b/backend/src/services/file_watcher_service.py @@ -3,6 +3,7 @@ """ import os import time +import logging import threading from pathlib import Path from typing import Optional, Callable @@ -13,6 +14,7 @@ from ..core.config import get_settings from ..core.database import get_db from .knowledge_base_service import KnowledgeBaseService +logger = logging.getLogger(__name__) settings = get_settings() @@ -27,19 +29,19 @@ class KnowledgeBaseHandler(FileSystemEventHandler): def on_created(self, event): """处理文件创建事件""" if not event.is_directory and self._is_supported_file(event.src_path): - print(f"检测到新文件: {event.src_path}") + logger.info(f"检测到新文件: {event.src_path}") self._process_file_async(event.src_path, "created") def on_modified(self, event): """处理文件修改事件""" if not event.is_directory and self._is_supported_file(event.src_path): - print(f"检测到文件修改: {event.src_path}") + logger.info(f"检测到文件修改: {event.src_path}") self._process_file_async(event.src_path, "modified") def on_deleted(self, event): """处理文件删除事件""" if not event.is_directory and self._is_supported_file(event.src_path): - print(f"检测到文件删除: {event.src_path}") + logger.info(f"检测到文件删除: {event.src_path}") self._handle_file_deletion(event.src_path) def _is_supported_file(self, file_path: str) -> bool: @@ -54,10 +56,10 @@ class KnowledgeBaseHandler(FileSystemEventHandler): time.sleep(1) result = self.kb_service.process_file(file_path) - print(f"文件处理结果 ({event_type}): {result}") + logger.info(f"文件处理结果 ({event_type}): {result}") except Exception as e: - print(f"处理文件失败: {file_path}, 错误: {str(e)}") + logger.error(f"处理文件失败: {file_path}, 错误: {str(e)}") # 在后台线程中处理 thread = threading.Thread(target=process) @@ -80,12 +82,12 @@ class KnowledgeBaseHandler(FileSystemEventHandler): if document: result = self.kb_service.delete_document(document.id) - print(f"文件删除处理结果: {result}") + logger.info(f"文件删除处理结果: {result}") else: - print(f"未找到对应的数据库记录: {file_path}") + logger.info(f"未找到对应的数据库记录: {file_path}") except Exception as e: - print(f"处理文件删除失败: {file_path}, 错误: {str(e)}") + logger.error(f"处理文件删除失败: {file_path}, 错误: {str(e)}") class FileWatcherService: @@ -99,15 +101,15 @@ class FileWatcherService: def start(self): """启动文件监控""" if self.is_running: - print("文件监控服务已在运行") + logger.info("文件监控服务已在运行") return if not settings.enable_file_watcher: - print("文件监控服务已禁用") + logger.info("文件监控服务已禁用") return if not self.knowledge_base_dir.exists(): - print(f"知识库目录不存在: {self.knowledge_base_dir}") + logger.info(f"知识库目录不存在: {self.knowledge_base_dir}") return try: @@ -130,19 +132,19 @@ class FileWatcherService: self.observer.start() self.is_running = True - print(f"文件监控服务已启动,监控目录: {self.knowledge_base_dir}") + logger.info(f"文件监控服务已启动,监控目录: {self.knowledge_base_dir}") # 确保子目录对应的系统知识库存在 - print("确保系统知识库与目录同步...") + logger.info("确保系统知识库与目录同步...") kb_service.ensure_system_knowledge_bases() # 执行初始扫描 - print("执行初始知识库扫描...") + logger.info("执行初始知识库扫描...") scan_result = kb_service.scan_directory() - print(f"初始扫描结果: {scan_result}") + logger.info(f"初始扫描结果: {scan_result}") except Exception as e: - print(f"启动文件监控服务失败: {str(e)}") + logger.error(f"启动文件监控服务失败: {str(e)}") self.is_running = False def stop(self): @@ -151,7 +153,7 @@ class FileWatcherService: self.observer.stop() self.observer.join() self.is_running = False - print("文件监控服务已停止") + logger.info("文件监控服务已停止") def is_active(self) -> bool: """检查监控服务是否活跃""" diff --git a/backend/src/services/image_generation_service.py b/backend/src/services/image_generation_service.py index 3534f88..9babfaf 100644 --- a/backend/src/services/image_generation_service.py +++ b/backend/src/services/image_generation_service.py @@ -4,12 +4,15 @@ """ import base64 +import logging import uuid from pathlib import Path from typing import Dict, Any import httpx from src.core.config import settings +logger = logging.getLogger(__name__) + class ImageGenerationService: """统一的图像生成服务基类""" @@ -21,9 +24,9 @@ class ImageGenerationService: async def call_api(self, model: str, payload: dict) -> dict: """调用硅基流动 API""" - print(f"[DEBUG] 调用SiliconFlow API: {self.base_url}") - print(f"[DEBUG] 模型: {model}") - print(f"[DEBUG] 请求参数: {payload}") + logger.debug(f"调用SiliconFlow API: {self.base_url}") + logger.debug(f"模型: {model}") + logger.debug(f"请求参数: {payload}") async with httpx.AsyncClient(timeout=self.timeout) as client: response = await client.post( @@ -37,9 +40,9 @@ class ImageGenerationService: response.raise_for_status() result = response.json() - print(f"[DEBUG] API响应状态: {response.status_code}") - print(f"[DEBUG] API响应类型: {type(result)}") - print(f"[DEBUG] API响应键: {list(result.keys()) if isinstance(result, dict) else '非字典类型'}") + logger.debug(f"API响应状态: {response.status_code}") + logger.debug(f"API响应类型: {type(result)}") + logger.debug(f"API响应键: {list(result.keys()) if isinstance(result, dict) else '非字典类型'}") return result @@ -53,7 +56,7 @@ class ImageGenerationService: # 移除可能的空白字符 base64_data = base64_data.strip() - print(f"[DEBUG] 开始解码base64数据,长度: {len(base64_data)}") + logger.debug(f"开始解码base64数据,长度: {len(base64_data)}") # 解码 base64 数据 image_bytes = base64.b64decode(base64_data) @@ -62,7 +65,7 @@ class ImageGenerationService: if len(image_bytes) == 0: raise Exception("解码后的图像数据为空") - print(f"[DEBUG] 图像数据大小: {len(image_bytes)} bytes") + logger.debug(f"图像数据大小: {len(image_bytes)} bytes") # 确保目录存在 image_dir = Path(settings.generated_images_dir) @@ -73,13 +76,13 @@ class ImageGenerationService: with open(image_path, "wb") as f: f.write(image_bytes) - print(f"[DEBUG] 图像已保存到: {image_path}") + logger.debug(f"图像已保存到: {image_path}") return str(image_path) except Exception as e: - print(f"[ERROR] 保存图像失败: {str(e)}") - print(f"[ERROR] base64数据长度: {len(base64_data) if base64_data else 0}") - print(f"[ERROR] base64数据前100字符: {base64_data[:100] if base64_data else 'None'}") + logger.error(f"保存图像失败: {str(e)}") + logger.error(f"base64数据长度: {len(base64_data) if base64_data else 0}") + logger.error(f"base64数据前100字符: {base64_data[:100] if base64_data else 'None'}") raise Exception(f"保存图像失败: {str(e)}") def generate_image_id(self) -> str: diff --git a/backend/src/services/image_to_image_service.py b/backend/src/services/image_to_image_service.py index c3570f9..e7de400 100644 --- a/backend/src/services/image_to_image_service.py +++ b/backend/src/services/image_to_image_service.py @@ -4,12 +4,15 @@ """ import base64 +import logging import httpx from typing import List, Dict, Any, Optional from pathlib import Path from fastapi import UploadFile from .image_generation_service import ImageGenerationService +logger = logging.getLogger(__name__) + class ImageToImageService(ImageGenerationService): """图生图服务""" @@ -59,70 +62,70 @@ class ImageToImageService(ImageGenerationService): result = await self.call_api(self.MODEL, payload) # 调试信息 - print(f"API响应结构: {type(result)}") + logger.debug(f"API响应结构: {type(result)}") if isinstance(result, dict): - print(f"响应键: {list(result.keys())}") + logger.debug(f"响应键: {list(result.keys())}") if "images" in result: - print(f"图像数量: {len(result['images'])}") + logger.debug(f"图像数量: {len(result['images'])}") if result["images"]: - print(f"第一个图像键: {list(result['images'][0].keys()) if isinstance(result['images'][0], dict) else '非字典类型'}") - + logger.debug(f"第一个图像键: {list(result['images'][0].keys()) if isinstance(result['images'][0], dict) else '非字典类型'}") + # 保存结果 image_id = self.generate_image_id() edited_image_data = result["images"][0] - + # 检查是否为字典类型 if isinstance(edited_image_data, dict): - print(f"[DEBUG] 图像数据是字典,键: {list(edited_image_data.keys())}") - + logger.debug(f"图像数据是字典,键: {list(edited_image_data.keys())}") + # 优先检查URL字段 if "url" in edited_image_data: image_url = edited_image_data["url"] - print(f"[DEBUG] 检测到URL字段: {image_url}") + logger.info(f"检测到URL字段: {image_url}") try: image_path = await self.download_image_from_url(image_id, image_url) - print(f"[DEBUG] URL下载成功: {image_path}") + logger.info(f"URL下载成功: {image_path}") except Exception as e: - print(f"[ERROR] URL下载失败: {str(e)}") + logger.error(f"URL下载失败: {str(e)}") raise Exception(f"下载图像失败: {str(e)}") - + # 如果没有URL,尝试base64字段 elif any(key in edited_image_data for key in ["b64_json", "b64", "data"]): image_b64 = edited_image_data.get("b64_json") or edited_image_data.get("b64") or edited_image_data.get("data") - print(f"[DEBUG] 检测到base64字段,长度: {len(image_b64) if image_b64 else 0}") + logger.debug(f"检测到base64字段,长度: {len(image_b64) if image_b64 else 0}") try: image_path = self.save_image(image_id, image_b64) - print(f"[DEBUG] base64保存成功: {image_path}") + logger.info(f"base64保存成功: {image_path}") except Exception as e: - print(f"[ERROR] base64保存失败: {str(e)}") + logger.error(f"base64保存失败: {str(e)}") raise Exception(f"保存图像失败: {str(e)}") - + else: - print(f"[ERROR] 未找到有效的图像数据字段") + logger.error(f"未找到有效的图像数据字段") raise Exception("API响应中未找到有效的图像数据") - + # 如果直接是字符串,判断是URL还是base64 elif isinstance(edited_image_data, str): - print(f"[DEBUG] 图像数据是字符串,长度: {len(edited_image_data)}") + logger.debug(f"图像数据是字符串,长度: {len(edited_image_data)}") if edited_image_data.startswith('http'): - print(f"[DEBUG] 检测到URL字符串: {edited_image_data}") + logger.info(f"检测到URL字符串: {edited_image_data}") try: image_path = await self.download_image_from_url(image_id, edited_image_data) - print(f"[DEBUG] URL下载成功: {image_path}") + logger.info(f"URL下载成功: {image_path}") except Exception as e: - print(f"[ERROR] URL下载失败: {str(e)}") + logger.error(f"URL下载失败: {str(e)}") raise Exception(f"下载图像失败: {str(e)}") else: - print(f"[DEBUG] 检测到base64字符串") + logger.debug(f"检测到base64字符串") try: image_path = self.save_image(image_id, edited_image_data) - print(f"[DEBUG] base64保存成功: {image_path}") + logger.info(f"base64保存成功: {image_path}") except Exception as e: - print(f"[ERROR] base64保存失败: {str(e)}") + logger.error(f"base64保存失败: {str(e)}") raise Exception(f"保存图像失败: {str(e)}") - + else: - print(f"[ERROR] 未知的图像数据类型: {type(edited_image_data)}") + logger.error(f"未知的图像数据类型: {type(edited_image_data)}") raise Exception(f"不支持的图像数据类型: {type(edited_image_data)}") return { @@ -165,70 +168,70 @@ class ImageToImageService(ImageGenerationService): result = await self.call_api(self.MODEL, payload) # 调试信息 - print(f"API响应结构: {type(result)}") + logger.debug(f"API响应结构: {type(result)}") if isinstance(result, dict): - print(f"响应键: {list(result.keys())}") + logger.debug(f"响应键: {list(result.keys())}") if "images" in result: - print(f"图像数量: {len(result['images'])}") - + logger.debug(f"图像数量: {len(result['images'])}") + # 处理变体结果 variations = [] for i, img_data in enumerate(result["images"]): - print(f"[DEBUG] 处理第 {i+1} 个变体") + logger.debug(f"处理第 {i+1} 个变体") image_id = self.generate_image_id() - + # 检查是否为字典类型 if isinstance(img_data, dict): - print(f"[DEBUG] 图像数据是字典,键: {list(img_data.keys())}") - + logger.debug(f"图像数据是字典,键: {list(img_data.keys())}") + # 优先检查URL字段 if "url" in img_data: image_url = img_data["url"] - print(f"[DEBUG] 检测到URL字段: {image_url}") + logger.info(f"检测到URL字段: {image_url}") try: image_path = await self.download_image_from_url(image_id, image_url) - print(f"[DEBUG] URL下载成功: {image_path}") + logger.info(f"URL下载成功: {image_path}") except Exception as e: - print(f"[ERROR] URL下载失败: {str(e)}") + logger.error(f"URL下载失败: {str(e)}") raise Exception(f"下载图像失败: {str(e)}") - + # 如果没有URL,尝试base64字段 elif any(key in img_data for key in ["b64_json", "b64", "data"]): image_b64 = img_data.get("b64_json") or img_data.get("b64") or img_data.get("data") - print(f"[DEBUG] 检测到base64字段,长度: {len(image_b64) if image_b64 else 0}") + logger.debug(f"检测到base64字段,长度: {len(image_b64) if image_b64 else 0}") try: image_path = self.save_image(image_id, image_b64) - print(f"[DEBUG] base64保存成功: {image_path}") + logger.info(f"base64保存成功: {image_path}") except Exception as e: - print(f"[ERROR] base64保存失败: {str(e)}") + logger.error(f"base64保存失败: {str(e)}") raise Exception(f"保存图像失败: {str(e)}") - + else: - print(f"[ERROR] 未找到有效的图像数据字段") + logger.error(f"未找到有效的图像数据字段") raise Exception("API响应中未找到有效的图像数据") - + # 如果直接是字符串,判断是URL还是base64 elif isinstance(img_data, str): - print(f"[DEBUG] 图像数据是字符串,长度: {len(img_data)}") + logger.debug(f"图像数据是字符串,长度: {len(img_data)}") if img_data.startswith('http'): - print(f"[DEBUG] 检测到URL字符串: {img_data}") + logger.info(f"检测到URL字符串: {img_data}") try: image_path = await self.download_image_from_url(image_id, img_data) - print(f"[DEBUG] URL下载成功: {image_path}") + logger.info(f"URL下载成功: {image_path}") except Exception as e: - print(f"[ERROR] URL下载失败: {str(e)}") + logger.error(f"URL下载失败: {str(e)}") raise Exception(f"下载图像失败: {str(e)}") else: - print(f"[DEBUG] 检测到base64字符串") + logger.debug(f"检测到base64字符串") try: image_path = self.save_image(image_id, img_data) - print(f"[DEBUG] base64保存成功: {image_path}") + logger.info(f"base64保存成功: {image_path}") except Exception as e: - print(f"[ERROR] base64保存失败: {str(e)}") + logger.error(f"base64保存失败: {str(e)}") raise Exception(f"保存图像失败: {str(e)}") - + else: - print(f"[ERROR] 未知的图像数据类型: {type(img_data)}") + logger.error(f"未知的图像数据类型: {type(img_data)}") raise Exception(f"不支持的图像数据类型: {type(img_data)}") variations.append({ @@ -320,7 +323,7 @@ class ImageToImageService(ImageGenerationService): async def download_image_from_url(self, image_id: str, image_url: str) -> str: """从URL下载图像""" try: - print(f"[DEBUG] 下载图像URL: {image_url}") + logger.debug(f"下载图像URL: {image_url}") # 下载图像 async with httpx.AsyncClient(timeout=30) as client: @@ -331,7 +334,7 @@ class ImageToImageService(ImageGenerationService): if len(image_bytes) == 0: raise Exception("下载的图像数据为空") - print(f"[DEBUG] 下载图像数据大小: {len(image_bytes)} bytes") + logger.debug(f"下载图像数据大小: {len(image_bytes)} bytes") # 确保目录存在 image_dir = Path(self.get_image_dir()) @@ -342,11 +345,11 @@ class ImageToImageService(ImageGenerationService): with open(image_path, "wb") as f: f.write(image_bytes) - print(f"[DEBUG] 图像已保存到: {image_path}") + logger.debug(f"图像已保存到: {image_path}") return str(image_path) except Exception as e: - print(f"[ERROR] 下载图像失败: {str(e)}") + logger.error(f"下载图像失败: {str(e)}") raise Exception(f"下载图像失败: {str(e)}") def get_image_dir(self) -> str: diff --git a/backend/src/services/knowledge_base_service.py b/backend/src/services/knowledge_base_service.py index b4adc08..c1a87ec 100644 --- a/backend/src/services/knowledge_base_service.py +++ b/backend/src/services/knowledge_base_service.py @@ -1,6 +1,7 @@ """ 知识库管理服务 """ +import logging import os import hashlib from pathlib import Path @@ -15,6 +16,8 @@ from ..models.user import User from ..core.config import get_settings from .document_service import DocumentService +logger = logging.getLogger(__name__) + settings = get_settings() @@ -106,7 +109,7 @@ class KnowledgeBaseService: existing.user_id = admin.id self.db.commit() self.db.refresh(existing) - print(f"[KB] 升级为系统知识库: {name} (id={existing.id})") + logger.info(f"升级为系统知识库: {name} (id={existing.id})") return existing.id # 找到 admin 用户(或任意 superuser)作为 owner @@ -123,7 +126,7 @@ class KnowledgeBaseService: self.db.add(kb) self.db.commit() self.db.refresh(kb) - print(f"[KB] 自动创建系统知识库: {name} (id={kb.id})") + logger.info(f"自动创建系统知识库: {name} (id={kb.id})") return kb.id def _resolve_knowledge_base(self, file_path: Path) -> Optional[int]: @@ -249,13 +252,13 @@ class KnowledgeBaseService: loop.run_until_complete(self.document_service.process_document(document.id)) loop.close() except Exception as e: - print(f"后台处理文档 {document.id} 失败: {e}") + logger.error(f"后台处理文档 {document.id} 失败: {e}") thread = threading.Thread(target=process_in_background, daemon=True) thread.start() except Exception as e: - print(f"启动文档处理线程失败: {e}") - + logger.error(f"启动文档处理线程失败: {e}") + return {"status": "new", "document_id": document.id, "message": "文档创建成功,正在处理中"} except Exception as e: @@ -287,13 +290,13 @@ class KnowledgeBaseService: loop.run_until_complete(self.document_service.process_document(document.id)) loop.close() except Exception as e: - print(f"后台处理文档 {document.id} 失败: {e}") + logger.error(f"后台处理文档 {document.id} 失败: {e}") thread = threading.Thread(target=process_in_background, daemon=True) thread.start() except Exception as e: - print(f"启动文档处理线程失败: {e}") - + logger.error(f"启动文档处理线程失败: {e}") + return {"status": "updated", "document_id": document.id, "message": "文档更新成功,正在重新处理中"} except Exception as e: @@ -482,7 +485,7 @@ class KnowledgeBaseService: all_files = [f for f in directory.rglob("*") if f.is_file() and self._is_supported_file(f)] total_files = len(all_files) - print(f" 找到 {total_files} 个支持的文件,开始处理...") + logger.info(f"找到 {total_files} 个支持的文件,开始处理...") for idx, file_path in enumerate(all_files, 1): try: @@ -506,7 +509,7 @@ class KnowledgeBaseService: # 显示进度(每10个文件或最后一个文件时显示) if idx % 10 == 0 or idx == total_files: percentage = (idx * 100) // total_files if total_files > 0 else 0 - print(f"\r 处理进度: {idx}/{total_files} ({percentage}%)", end="", flush=True) + logger.info(f"处理进度: {idx}/{total_files} ({percentage}%)") if existing_doc: # 检查是否需要更新 @@ -549,6 +552,6 @@ class KnowledgeBaseService: "error": str(e) }) - print() # 换行 + logger.info("文件处理完成") results["success"] = len(results["errors"]) == 0 return results diff --git a/backend/src/services/text_to_image_service.py b/backend/src/services/text_to_image_service.py index f9c28f1..5f06f4d 100644 --- a/backend/src/services/text_to_image_service.py +++ b/backend/src/services/text_to_image_service.py @@ -3,12 +3,15 @@ 支持 Kwai-Kolors/Kolors 和 Qwen/Qwen-Image 模型 """ +import logging import uuid import httpx from typing import List, Dict, Any from pathlib import Path from .image_generation_service import ImageGenerationService +logger = logging.getLogger(__name__) + class TextToImageService(ImageGenerationService): """文生图服务""" @@ -73,72 +76,72 @@ class TextToImageService(ImageGenerationService): result = await self.call_api(self.MODELS[model], payload) # 调试信息 - print(f"API响应结构: {type(result)}") + logger.debug(f"API响应结构: {type(result)}") if isinstance(result, dict): - print(f"响应键: {list(result.keys())}") + logger.debug(f"响应键: {list(result.keys())}") if "images" in result: - print(f"图像数量: {len(result['images'])}") + logger.debug(f"图像数量: {len(result['images'])}") if result["images"]: - print(f"第一个图像键: {list(result['images'][0].keys()) if isinstance(result['images'][0], dict) else '非字典类型'}") + logger.debug(f"第一个图像键: {list(result['images'][0].keys()) if isinstance(result['images'][0], dict) else '非字典类型'}") # 处理结果 images = [] for i, img_data in enumerate(result.get("images", [])): - print(f"[DEBUG] 处理第 {i+1} 张图像") + logger.debug(f"处理第 {i+1} 张图像") image_id = self.generate_image_id() # 检查是否为字典类型 if isinstance(img_data, dict): - print(f"[DEBUG] 图像数据是字典,键: {list(img_data.keys())}") - + logger.debug(f"图像数据是字典,键: {list(img_data.keys())}") + # 优先检查URL字段 if "url" in img_data: image_url = img_data["url"] - print(f"[DEBUG] 检测到URL字段: {image_url}") + logger.info(f"检测到URL字段: {image_url}") try: image_path = await self.download_image_from_url(image_id, image_url) - print(f"[DEBUG] URL下载成功: {image_path}") + logger.info(f"URL下载成功: {image_path}") except Exception as e: - print(f"[ERROR] URL下载失败: {str(e)}") + logger.error(f"URL下载失败: {str(e)}") raise Exception(f"下载图像失败: {str(e)}") - + # 如果没有URL,尝试base64字段 elif any(key in img_data for key in ["b64_json", "b64", "data"]): image_b64 = img_data.get("b64_json") or img_data.get("b64") or img_data.get("data") - print(f"[DEBUG] 检测到base64字段,长度: {len(image_b64) if image_b64 else 0}") + logger.debug(f"检测到base64字段,长度: {len(image_b64) if image_b64 else 0}") try: image_path = self.save_image(image_id, image_b64) - print(f"[DEBUG] base64保存成功: {image_path}") + logger.info(f"base64保存成功: {image_path}") except Exception as e: - print(f"[ERROR] base64保存失败: {str(e)}") + logger.error(f"base64保存失败: {str(e)}") raise Exception(f"保存图像失败: {str(e)}") - + else: - print(f"[ERROR] 未找到有效的图像数据字段") + logger.error(f"未找到有效的图像数据字段") raise Exception("API响应中未找到有效的图像数据") - + # 如果直接是字符串,判断是URL还是base64 elif isinstance(img_data, str): - print(f"[DEBUG] 图像数据是字符串,长度: {len(img_data)}") + logger.debug(f"图像数据是字符串,长度: {len(img_data)}") if img_data.startswith('http'): - print(f"[DEBUG] 检测到URL字符串: {img_data}") + logger.info(f"检测到URL字符串: {img_data}") try: image_path = await self.download_image_from_url(image_id, img_data) - print(f"[DEBUG] URL下载成功: {image_path}") + logger.info(f"URL下载成功: {image_path}") except Exception as e: - print(f"[ERROR] URL下载失败: {str(e)}") + logger.error(f"URL下载失败: {str(e)}") raise Exception(f"下载图像失败: {str(e)}") else: - print(f"[DEBUG] 检测到base64字符串") + logger.debug(f"检测到base64字符串") try: image_path = self.save_image(image_id, img_data) - print(f"[DEBUG] base64保存成功: {image_path}") + logger.info(f"base64保存成功: {image_path}") except Exception as e: - print(f"[ERROR] base64保存失败: {str(e)}") + logger.error(f"base64保存失败: {str(e)}") raise Exception(f"保存图像失败: {str(e)}") - + else: - print(f"[ERROR] 未知的图像数据类型: {type(img_data)}") + logger.error(f"未知的图像数据类型: {type(img_data)}") raise Exception(f"不支持的图像数据类型: {type(img_data)}") images.append({ @@ -182,7 +185,7 @@ class TextToImageService(ImageGenerationService): async def download_image_from_url(self, image_id: str, image_url: str) -> str: """从URL下载图像""" try: - print(f"[DEBUG] 下载图像URL: {image_url}") + logger.debug(f"下载图像URL: {image_url}") # 下载图像 async with httpx.AsyncClient(timeout=30) as client: @@ -193,7 +196,7 @@ class TextToImageService(ImageGenerationService): if len(image_bytes) == 0: raise Exception("下载的图像数据为空") - print(f"[DEBUG] 下载图像数据大小: {len(image_bytes)} bytes") + logger.debug(f"下载图像数据大小: {len(image_bytes)} bytes") # 确保目录存在 image_dir = Path(self.get_image_dir()) @@ -204,11 +207,11 @@ class TextToImageService(ImageGenerationService): with open(image_path, "wb") as f: f.write(image_bytes) - print(f"[DEBUG] 图像已保存到: {image_path}") + logger.debug(f"图像已保存到: {image_path}") return str(image_path) except Exception as e: - print(f"[ERROR] 下载图像失败: {str(e)}") + logger.error(f"下载图像失败: {str(e)}") raise Exception(f"下载图像失败: {str(e)}") def get_image_dir(self) -> str: