foramt: 规范化代码

This commit is contained in:
Wenjie Zhang 2026-03-17 19:46:09 +08:00
parent 90bfd20d8d
commit c7b2bd8d71
47 changed files with 162 additions and 4529 deletions

View File

@ -35,7 +35,7 @@ assignees: ''
# 例如
docker compose up -d
# 或
make start
make up
```

View File

@ -1,27 +1,16 @@
.PHONY: start stop logs lint format format_diff router-tests
.PHONY: up down logs lint format format_diff router-tests
PYTEST_ARGS ?=
pull:
bash docker/pull_image.sh python:3.12-slim
bash docker/pull_image.sh node:20-slim
bash docker/pull_image.sh node:20-alpine
bash docker/pull_image.sh milvusdb/milvus:v2.5.6
bash docker/pull_image.sh neo4j:5.26
bash docker/pull_image.sh minio/minio:RELEASE.2023-03-20T20-16-18Z
bash docker/pull_image.sh ghcr.io/astral-sh/uv:0.7.2
bash docker/pull_image.sh nginx:alpine
bash docker/pull_image.sh quay.io/coreos/etcd:v3.5.5
start:
up:
@if [ ! -f .env ]; then \
echo "Error: .env file not found. Please create it from .env.template"; \
exit 1; \
fi
docker compose up -d
stop:
down:
docker compose down
logs:
@ -35,15 +24,15 @@ logs:
######################
lint:
uv run python -m ruff check .
uv run python -m ruff format --check src
uv run python -m ruff check --select I src
uv run python -m ruff check backend/package
uv run python -m ruff format --check backend/package
uv run python -m ruff check --select I backend/package
format:
uv run python -m ruff format .
uv run python -m ruff check . --fix
uv run python -m ruff check --select I src --fix
cd web && npm run format
uv run python -m ruff format backend/package
uv run python -m ruff check backend/package --fix
uv run python -m ruff check --select I backend/package --fix
docker compose exec -T web pnpm run format
router-tests:
docker compose exec -T api uv run --group test pytest test/api $(PYTEST_ARGS)

View File

@ -1,99 +1,4 @@
import asyncio
import importlib
import inspect
from pathlib import Path
from server.utils.singleton import SingletonMeta
from yuxi.agents.common import BaseAgent
from yuxi.utils import logger
class AgentManager(metaclass=SingletonMeta):
def __init__(self):
self._classes = {}
self._instances = {} # 存储已创建的 agent 实例
def register_agent(self, agent_class):
self._classes[agent_class.__name__] = agent_class
def init_all_agents(self):
for agent_id in self._classes.keys():
self.get_agent(agent_id)
def get_agent(self, agent_id, reload=False, reload_graph=False, **kwargs):
# 检查是否已经创建了该 agent 的实例
if reload or agent_id not in self._instances:
agent_class = self._classes[agent_id]
self._instances[agent_id] = agent_class()
# 如果仅需要重新加载 graph则清空 graph 缓存
if reload_graph and agent_id in self._instances:
self._instances[agent_id].reload_graph()
return self._instances[agent_id]
def get_agents(self):
return list(self._instances.values())
async def reload_all(self):
for agent_id in self._classes.keys():
self.get_agent(agent_id, reload=True)
async def get_agents_info(self, include_configurable_items: bool = True):
agents = self.get_agents()
return await asyncio.gather(
*[a.get_info(include_configurable_items=include_configurable_items) for a in agents]
)
def auto_discover_agents(self):
"""自动发现并注册 yuxi/agents/ 下的所有智能体。
遍历 yuxi/agents/ 目录下的所有子文件夹如果子文件夹包含 __init__.py
则尝试从中导入 BaseAgent 的子类并注册(使用自动导入的方式支持私有agent)
"""
# 获取 agents 目录的路径
agents_dir = Path(__file__).parent
# 遍历所有子目录
for item in agents_dir.iterdir():
# logger.info(f"尝试导入模块:{item}")
# 跳过非目录、common 目录、__pycache__ 等
if not item.is_dir() or item.name.startswith("_") or item.name in {"common", "skills"}:
continue
# 检查是否有 __init__.py 文件
init_file = item / "__init__.py"
if not init_file.exists():
logger.warning(f"{item} 不是一个有效的模块")
continue
# 尝试导入模块
try:
module_name = f"yuxi.agents.{item.name}"
module = importlib.import_module(module_name)
# 查找模块中所有 BaseAgent 的子类
for name, obj in inspect.getmembers(module):
if (
inspect.isclass(obj)
and issubclass(obj, BaseAgent)
and obj is not BaseAgent
and obj.__module__.startswith(module_name)
):
logger.info(f"自动发现智能体: {obj.__name__} 来自 {item.name}")
self.register_agent(obj)
except Exception as e:
logger.warning(f"无法从 {item.name} 加载智能体: {e}")
agent_manager = AgentManager()
# 自动发现并注册所有智能体
agent_manager.auto_discover_agents()
agent_manager.init_all_agents()
# 从 buildin 模块导入 agent_manager
from yuxi.agents.buildin import agent_manager
__all__ = ["agent_manager"]
if __name__ == "__main__":
pass

View File

@ -0,0 +1,99 @@
import asyncio
import importlib
import inspect
from pathlib import Path
from server.utils.singleton import SingletonMeta
from yuxi.agents.common import BaseAgent
from yuxi.utils import logger
class AgentManager(metaclass=SingletonMeta):
def __init__(self):
self._classes = {}
self._instances = {} # 存储已创建的 agent 实例
def register_agent(self, agent_class):
self._classes[agent_class.__name__] = agent_class
def init_all_agents(self):
for agent_id in self._classes.keys():
self.get_agent(agent_id)
def get_agent(self, agent_id, reload=False, reload_graph=False, **kwargs):
# 检查是否已经创建了该 agent 的实例
if reload or agent_id not in self._instances:
agent_class = self._classes[agent_id]
self._instances[agent_id] = agent_class()
# 如果仅需要重新加载 graph则清空 graph 缓存
if reload_graph and agent_id in self._instances:
self._instances[agent_id].reload_graph()
return self._instances[agent_id]
def get_agents(self):
return list(self._instances.values())
async def reload_all(self):
for agent_id in self._classes.keys():
self.get_agent(agent_id, reload=True)
async def get_agents_info(self, include_configurable_items: bool = True):
agents = self.get_agents()
return await asyncio.gather(
*[a.get_info(include_configurable_items=include_configurable_items) for a in agents]
)
def auto_discover_agents(self):
"""自动发现并注册 yuxi/agents/buildin/ 下的所有智能体。
遍历 yuxi/agents/buildin/ 目录下的所有子文件夹如果子文件夹包含 __init__.py
则尝试从中导入 BaseAgent 的子类并注册(使用自动导入的方式支持私有agent)
"""
# 获取 agents 目录的路径
agents_dir = Path(__file__).parent
# 遍历所有子目录
for item in agents_dir.iterdir():
# logger.info(f"尝试导入模块:{item}")
# 跳过非目录、common 目录、__pycache__ 等
if not item.is_dir() or item.name.startswith("_"):
continue
# 检查是否有 __init__.py 文件
init_file = item / "__init__.py"
if not init_file.exists():
logger.warning(f"{item} 不是一个有效的模块")
continue
# 尝试导入模块
try:
module_name = f"yuxi.agents.buildin.{item.name}"
module = importlib.import_module(module_name)
# 查找模块中所有 BaseAgent 的子类
for name, obj in inspect.getmembers(module):
if (
inspect.isclass(obj)
and issubclass(obj, BaseAgent)
and obj is not BaseAgent
and obj.__module__.startswith(module_name)
):
logger.info(f"自动发现智能体: {obj.__name__} 来自 {item.name}")
self.register_agent(obj)
except Exception as e:
logger.warning(f"无法从 {item.name} 加载智能体: {e}")
agent_manager = AgentManager()
# 自动发现并注册所有智能体
agent_manager.auto_discover_agents()
agent_manager.init_all_agents()
__all__ = ["agent_manager"]
if __name__ == "__main__":
pass

View File

@ -10,7 +10,7 @@ from yuxi import config, graph_base
from yuxi.agents.common.toolkits.registry import ToolExtraMetadata, _all_tool_instances, _extra_registry, tool
from yuxi.storage.minio import aupload_file_to_minio
from yuxi.utils import logger
from yuxi.utils.question_utils import normalize_questions, normalize_options
from yuxi.utils.question_utils import normalize_questions
# Lazy initialization for TavilySearch (only when API key is available)
_tavily_search_instance = None

View File

@ -11,7 +11,6 @@ from collections.abc import AsyncIterator
from fastapi import HTTPException
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession
from yuxi.agents import agent_manager
from yuxi.repositories.agent_run_repository import TERMINAL_RUN_STATUSES, AgentRunRepository
from yuxi.repositories.conversation_repository import ConversationRepository
@ -230,7 +229,6 @@ async def stream_agent_run_events(
async def get_active_run_by_thread(*, thread_id: str, current_user_id: str, db: AsyncSession) -> dict:
from sqlalchemy import select
from yuxi.storage.postgres.models_business import AgentRun
result = await db.execute(

View File

@ -8,7 +8,6 @@ from typing import Any
from langchain.messages import AIMessage, AIMessageChunk, HumanMessage
from langgraph.types import Command
from yuxi import config as conf
from yuxi.agents import agent_manager
from yuxi.plugins.guard import content_guard
@ -17,9 +16,10 @@ from yuxi.repositories.conversation_repository import ConversationRepository
from yuxi.storage.postgres.manager import pg_manager
from yuxi.utils.logging_config import logger
from yuxi.utils.question_utils import (
normalize_questions as _normalize_interrupt_questions,
normalize_options as _normalize_interrupt_options,
normalize_legacy_question,
)
from yuxi.utils.question_utils import (
normalize_questions as _normalize_interrupt_questions,
)

View File

@ -3,7 +3,6 @@ from datetime import UTC, datetime
from fastapi import HTTPException, UploadFile
from sqlalchemy.ext.asyncio import AsyncSession
from yuxi.agents import agent_manager
from yuxi.repositories.conversation_repository import ConversationRepository
from yuxi.services.doc_converter import (

View File

@ -8,7 +8,6 @@ from pathlib import Path
import aiofiles
from fastapi import UploadFile
from yuxi.config import config as app_config
from yuxi.knowledge.indexing import process_file_to_markdown
from yuxi.utils import logger

View File

@ -3,7 +3,6 @@ import traceback
from fastapi import HTTPException
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from yuxi.storage.postgres.models_business import Conversation, Message, MessageFeedback
from yuxi.utils.logging_config import logger

View File

@ -1,6 +1,5 @@
from fastapi import HTTPException
from sqlalchemy.ext.asyncio import AsyncSession
from yuxi.agents import agent_manager
from yuxi.repositories.conversation_repository import ConversationRepository
from yuxi.utils.logging_config import logger

View File

@ -16,7 +16,6 @@ from typing import Any, cast
from langchain_mcp_adapters.client import MultiServerMCPClient
from sqlalchemy import func, select
from sqlalchemy.ext.asyncio import AsyncSession
from yuxi.storage.postgres.models_business import MCPServer
from yuxi.utils import logger

View File

@ -10,7 +10,6 @@ from dataclasses import dataclass, field
from sqlalchemy import select
from sqlalchemy.exc import OperationalError
from yuxi.repositories.agent_run_repository import TERMINAL_RUN_STATUSES, AgentRunRepository
from yuxi.services.chat_stream_service import stream_agent_chat
from yuxi.services.run_queue_service import (

View File

@ -11,7 +11,6 @@ from typing import Any
import yaml
from sqlalchemy.ext.asyncio import AsyncSession
from yuxi import config as sys_config
from yuxi.repositories.skill_repository import SkillRepository
from yuxi.services.mcp_service import get_mcp_server_names

View File

@ -11,10 +11,10 @@ from datetime import timedelta
from io import BytesIO
from urllib3 import BaseHTTPResponse
from yuxi.utils import logger
from minio import Minio
from minio.error import S3Error
from yuxi.utils import logger
class StorageError(Exception):

View File

@ -7,12 +7,12 @@ from contextlib import asynccontextmanager
from sqlalchemy import text
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
from sqlalchemy.orm import declarative_base
from server.utils.singleton import SingletonMeta
from yuxi.storage.postgres.models_business import Base as BusinessBase
from yuxi.storage.postgres.models_knowledge import Base as KnowledgeBase
from yuxi.utils import logger
from server.utils.singleton import SingletonMeta
# 合并两个 Base
CombinedBase = declarative_base()

View File

@ -17,7 +17,6 @@ from sqlalchemy import (
)
from sqlalchemy.ext.declarative import declarative_base
from sqlalchemy.orm import relationship
from yuxi.utils.datetime_utils import format_utc_datetime, utc_now_naive
Base = declarative_base()

View File

@ -14,7 +14,6 @@ from sqlalchemy import (
UniqueConstraint,
)
from sqlalchemy.dialects.postgresql import JSONB
from yuxi.storage.postgres.models_business import Base
from yuxi.utils.datetime_utils import utc_now_naive

View File

@ -1,4 +1,5 @@
"""问题和选项规范化工具"""
import uuid
from typing import Any
@ -21,9 +22,7 @@ def normalize_options(raw_options: Any) -> list[dict[str, str]]:
return options
def normalize_questions(
raw_questions: Any, default_question_id_prefix: str = "q"
) -> list[dict[str, Any]]:
def normalize_questions(raw_questions: Any, default_question_id_prefix: str = "q") -> list[dict[str, Any]]:
"""规范化问题列表"""
if not isinstance(raw_questions, list):
return []
@ -37,9 +36,7 @@ def normalize_questions(
if not question:
continue
question_id = str(
item.get("question_id") or f"{default_question_id_prefix}-{idx + 1}"
).strip()
question_id = str(item.get("question_id") or f"{default_question_id_prefix}-{idx + 1}").strip()
if not question_id:
question_id = str(uuid.uuid4())

View File

@ -1,525 +0,0 @@
import asyncio
import hashlib
import json
import pathlib
import httpx
import typer
from rich.console import Console
from rich.progress import BarColumn, Progress, SpinnerColumn, TextColumn, TimeElapsedColumn
app = typer.Typer()
console = Console()
async def login(client: httpx.AsyncClient, base_url: str, username: str, password: str) -> str | None:
"""Logs in to the API and returns the access token."""
try:
response = await client.post(
f"{base_url}/auth/token",
data={"username": username, "password": password},
)
response.raise_for_status()
return response.json().get("access_token")
except httpx.HTTPStatusError as e:
console.print(f"[bold red]Login failed: {e.response.status_code} - {e.response.text}[/bold red]")
return None
except httpx.RequestError as e:
console.print(f"[bold red]Login request failed: {e}[/bold red]")
return None
async def check_task_status(client: httpx.AsyncClient, base_url: str, task_id: str) -> str | None:
"""Check the status of a task. Returns status string or None if failed."""
try:
response = await client.get(f"{base_url}/tasks/{task_id}")
response.raise_for_status()
task_data = response.json().get("task", {})
return task_data.get("status")
except httpx.HTTPStatusError as e:
console.print(f"[bold yellow]Warning: Failed to check task {task_id}: {e.response.status_code}[/bold yellow]")
return None
except httpx.RequestError as e:
console.print(f"[bold yellow]Warning: Failed to check task {task_id}: {e}[/bold yellow]")
return None
async def upload_file(
client: httpx.AsyncClient,
base_url: str,
db_id: str,
file_path: pathlib.Path,
) -> str | None:
"""Uploads a single file and returns its server-side path."""
try:
with open(file_path, "rb") as f:
files = {"file": (file_path.name, f, "application/octet-stream")}
response = await client.post(
f"{base_url}/knowledge/files/upload",
params={"db_id": db_id},
files=files,
timeout=300, # 5 minutes timeout for large files
)
response.raise_for_status()
return response.json().get("file_path")
except httpx.HTTPStatusError as e:
console.print(
f"[bold red]Failed to upload {file_path.name}: {e.response.status_code} - {e.response.text}[/bold red]"
)
return None
except httpx.RequestError as e:
console.print(f"[bold red]Failed to upload {file_path.name}: {e}[/bold red]")
return None
async def process_document(
client: httpx.AsyncClient,
base_url: str,
db_id: str,
server_file_path: str,
enable_ocr: str = "paddlex_ocr",
chunk_size: int = 1000,
chunk_overlap: int = 200,
use_qa_split: bool = False,
qa_separator: str = "\n\n\n",
) -> tuple[bool, str | None]:
"""Triggers the processing of an uploaded file in the knowledge base."""
# Prepare processing parameters
params = {
"chunk_size": chunk_size,
"chunk_overlap": chunk_overlap,
"enable_ocr": enable_ocr,
"use_qa_split": use_qa_split,
"qa_separator": qa_separator,
"content_type": "file",
}
try:
response = await client.post(
f"{base_url}/knowledge/databases/{db_id}/documents",
json={"items": [server_file_path], "params": params},
timeout=600, # 10 minutes timeout for processing
)
response.raise_for_status()
result = response.json()
# Handle asynchronous ingest response
overall_status = result.get("status")
if overall_status == "queued":
task_id = result.get("task_id")
extra = f" (task id: {task_id})" if task_id else ""
console.print(
f"[bold cyan]Ingestion queued for {server_file_path}{extra}. "
"Track progress in the task center.[/bold cyan]"
)
return True, task_id
# Check if the overall request was successful for synchronous responses
if overall_status != "success":
console.print(
f"[bold yellow]Processing warning for {server_file_path}: {result.get('message')}[/bold yellow]"
)
return False, None
# Check the specific file's processing status in the items array
items = result.get("items", [])
if not items:
console.print(f"[bold red]No processing result for {server_file_path}[/bold red]")
return False, None
# Since we only sent one file, check the first item
item = items[0]
# Check for both 'success' and 'done' status (different APIs might use different status values)
if item.get("status") in ["success", "done"]:
return True, None
else:
# Get more detailed error information
error_msg = item.get("message", "")
error_detail = item.get("detail", "")
error_reason = item.get("reason", "")
# Combine all available error information
error_info = []
if error_msg:
error_info.append(error_msg)
if error_detail:
error_info.append(error_detail)
if error_reason:
error_info.append(error_reason)
if not error_info:
error_info = ["Unknown error"]
full_error = " | ".join(error_info)
console.print(f"[bold red]Failed to process {server_file_path}: {full_error}[/bold red]")
# Also log the full item for debugging
console.print(f"[dim]Debug - Full item response: {item}[/dim]")
return False, None
except httpx.HTTPStatusError as e:
console.print(
f"[bold red]Failed to process {server_file_path}: {e.response.status_code} - {e.response.text}[/bold red]"
)
return False, None
except httpx.RequestError as e:
console.print(f"[bold red]Failed to process {server_file_path}: {e}[/bold red]")
return False, None
async def upload_single_file(
client: httpx.AsyncClient,
base_url: str,
db_id: str,
file_path: pathlib.Path,
progress: Progress,
task_id: int,
) -> str | None:
"""Upload a single file and return server file path."""
server_file_path = await upload_file(client, base_url, db_id, file_path)
if server_file_path:
progress.update(task_id, advance=1, postfix=f"Uploaded {file_path.name}")
else:
progress.update(task_id, advance=1, postfix=f"Failed: {file_path.name}")
return server_file_path
async def add_batch_to_knowledge_base(
client: httpx.AsyncClient,
base_url: str,
db_id: str,
server_file_paths: list[str],
enable_ocr: str = "paddlex_ocr",
chunk_size: int = 1000,
chunk_overlap: int = 200,
use_qa_split: bool = False,
qa_separator: str = "\n\n\n",
) -> tuple[bool, str | None]:
"""Add a batch of files to knowledge base and return task_id."""
if not server_file_paths:
return True, None
# Prepare processing parameters
params = {
"chunk_size": chunk_size,
"chunk_overlap": chunk_overlap,
"enable_ocr": enable_ocr,
"use_qa_split": use_qa_split,
"qa_separator": qa_separator,
"content_type": "file",
}
try:
response = await client.post(
f"{base_url}/knowledge/databases/{db_id}/documents",
json={"items": server_file_paths, "params": params},
timeout=600, # 10 minutes timeout for processing
)
response.raise_for_status()
result = response.json()
overall_status = result.get("status")
if overall_status == "queued":
task_id = result.get("task_id")
extra = f" (task id: {task_id})" if task_id else ""
console.print(
f"[bold cyan]Batch of {len(server_file_paths)} files queued for processing{extra}. "
"Track progress in the task center.[/bold cyan]"
)
return True, task_id
elif overall_status == "success":
console.print(f"[bold green]Batch of {len(server_file_paths)} files processed successfully[/bold green]")
return True, None
else:
console.print(f"[bold yellow]Batch processing warning: {result.get('message')}[/bold yellow]")
return False, None
except httpx.HTTPStatusError as e:
console.print(f"[bold red]Failed to process batch: {e.response.status_code} - {e.response.text}[/bold red]")
return False, None
except httpx.RequestError as e:
console.print(f"[bold red]Failed to process batch: {e}[/bold red]")
return False, None
async def wait_for_tasks_completion(
client: httpx.AsyncClient,
base_url: str,
task_ids: list[str],
poll_interval: int = 5,
) -> dict[str, str]:
"""Wait for all tasks to complete and return their final statuses."""
if not task_ids:
return {}
console.print(f"[bold cyan]Waiting for {len(task_ids)} tasks to complete...[/bold cyan]")
pending_tasks = task_ids.copy()
completed_tasks = {}
while pending_tasks:
for task_id in pending_tasks.copy():
status = await check_task_status(client, base_url, task_id)
if status:
if status in ["success", "failed", "cancelled"]:
completed_tasks[task_id] = status
pending_tasks.remove(task_id)
console.print(f"[dim]Task {task_id} completed with status: {status}[/dim]")
if pending_tasks:
await asyncio.sleep(poll_interval)
console.print(f"[bold green]All {len(task_ids)} tasks completed[/bold green]")
return completed_tasks
def get_file_hash(file_path: pathlib.Path) -> str:
"""Calculate SHA256 hash of a file."""
hash_sha256 = hashlib.sha256()
with open(file_path, "rb") as f:
for chunk in iter(lambda: f.read(4096), b""):
hash_sha256.update(chunk)
return hash_sha256.hexdigest()
def load_processed_files(record_file: pathlib.Path) -> set[str]:
"""Load the set of processed file hashes from the record file."""
if not record_file.exists():
return set()
try:
with open(record_file) as f:
data = json.load(f)
return set(data.get("processed_files", []))
except (OSError, json.JSONDecodeError) as e:
console.print(f"[bold yellow]Warning: Could not load processed files record: {e}[/bold yellow]")
return set()
def save_processed_files(record_file: pathlib.Path, processed_files: set[str]):
"""Save the set of processed file hashes to the record file."""
# Ensure the directory exists
record_file.parent.mkdir(parents=True, exist_ok=True)
try:
with open(record_file, "w") as f:
json.dump({"processed_files": list(processed_files)}, f, indent=2)
except OSError as e:
console.print(f"[bold red]Error: Could not save processed files record: {e}[/bold red]")
@app.command()
def upload(
db_id: str = typer.Option(..., help="The ID of the knowledge base."),
directory: pathlib.Path = typer.Option(
..., help="The directory containing files to upload.", exists=True, file_okay=False
),
pattern: list[str] = typer.Option(
["*.md"],
help="The glob patterns for files to upload (e.g., '*.pdf', '**/*.txt'). Can be specified multiple times.",
),
base_url: str = typer.Option("http://127.0.0.1:5050/api", help="The base URL of the API server."),
username: str = typer.Option(..., help="Admin username for login."),
password: str = typer.Option(..., help="Admin password for login."),
recursive: bool = typer.Option(False, "--recursive", "-r", help="Search for files recursively in subdirectories."),
record_file: pathlib.Path = typer.Option(
"scripts/tmp/batch_processed_files.txt", help="File to store processed files record."
),
chunk_size: int = typer.Option(1000, help="Chunk size for document processing."),
chunk_overlap: int = typer.Option(200, help="Chunk overlap for document processing."),
enable_ocr: str = typer.Option(
"paddlex_ocr", help="OCR engine to use (onnx_rapid_ocr, mineru_ocr, mineru_official, paddlex_ocr, disable)."
),
use_qa_split: bool = typer.Option(False, help="Whether to use QA splitting."),
qa_separator: str = typer.Option("\n\n\n", help="Separator for QA splitting."),
batch_size: int = typer.Option(20, help="Number of files to process in each batch."),
wait_for_completion: bool = typer.Option(True, help="Whether to wait for tasks to complete before next batch."),
poll_interval: int = typer.Option(5, help="Polling interval in seconds for checking task status."),
):
"""
Batch upload and process files into a Yuxi knowledge base.
"""
console.print(f"[bold green]Starting batch upload for knowledge base: {db_id}[/bold green]")
# Load previously processed files
processed_files = load_processed_files(record_file)
console.print(f"Loaded {len(processed_files)} previously processed files from record.")
# Discover files from multiple patterns
glob_method = directory.rglob if recursive else directory.glob
all_files = []
for pat in pattern:
files_for_pat = list(glob_method(pat))
all_files.extend(files_for_pat)
# Remove duplicates
all_files = list(set(all_files))
if not all_files:
patterns_str = "', '".join(pattern)
console.print(
f"[bold yellow]No files found in '{directory}' matching patterns: '{patterns_str}'. Aborting.[/bold yellow]"
)
raise typer.Exit()
# 过滤掉macos的隐藏文件
all_files = [f for f in all_files if not f.name.startswith("._")]
# Filter out already processed files
files_to_upload = []
skipped_files = []
for file_path in all_files:
file_hash = get_file_hash(file_path)
if file_hash in processed_files:
skipped_files.append(file_path)
else:
files_to_upload.append((file_path, file_hash))
if not files_to_upload:
console.print(
f"[bold green]All {len(all_files)} files have already been processed. Nothing to do.[/bold green]"
)
raise typer.Exit()
console.print(f"Found {len(all_files)} total files:")
console.print(f" - [green]New files to process:[/green] {len(files_to_upload)}")
console.print(f" - [blue]Already processed (skipped):[/blue] {len(skipped_files)}")
async def run():
async with httpx.AsyncClient() as client:
# Login
token = await login(client, base_url, username, password)
if not token:
raise typer.Exit(code=1)
client.headers = {"Authorization": f"Bearer {token}"}
# Process files in batches: upload 20 -> process 20 -> wait -> repeat
total_processed_files = []
total_upload_failures = []
total_processing_failures = []
all_successful_hashes = set()
# Split all files into batches
for batch_num in range(0, len(files_to_upload), batch_size):
batch_files = files_to_upload[batch_num : batch_num + batch_size]
batch_start = batch_num + 1
batch_end = min(batch_num + batch_size, len(files_to_upload))
console.print(
f"\n[bold yellow]=== Batch {batch_start}-{batch_end} of {len(files_to_upload)} ===[/bold yellow]"
)
# Step 1: Upload this batch of files sequentially
console.print(f"[blue]Step 1: Uploading {len(batch_files)} files...[/blue]")
successful_uploads = []
batch_upload_failures = []
with Progress(
SpinnerColumn(),
TextColumn("[progress.description]{task.description}"),
BarColumn(),
TextColumn("[progress.percentage]{task.percentage:>3.0f}%"),
TimeElapsedColumn(),
TextColumn("{task.fields[postfix]}"),
console=console,
transient=True,
) as progress:
upload_task_id = progress.add_task(
f"Uploading batch {batch_start}-{batch_end}...", total=len(batch_files), postfix=""
)
for file_path, file_hash in batch_files:
server_file_path = await upload_single_file(
client, base_url, db_id, file_path, progress, upload_task_id
)
if server_file_path:
successful_uploads.append((file_path, file_hash, server_file_path))
all_successful_hashes.add(file_hash)
else:
batch_upload_failures.append(file_path)
# Step 2: Process this batch if uploads succeeded
if successful_uploads:
console.print(f"[green]Step 2: Processing {len(successful_uploads)} uploaded files...[/green]")
# Extract server file paths
server_file_paths = [item[2] for item in successful_uploads]
# Submit batch to knowledge base
success, task_id = await add_batch_to_knowledge_base(
client,
base_url,
db_id,
server_file_paths,
enable_ocr=enable_ocr,
chunk_size=chunk_size,
chunk_overlap=chunk_overlap,
use_qa_split=use_qa_split,
qa_separator=qa_separator,
)
if success:
total_processed_files.extend([item[0] for item in successful_uploads])
# Step 3: Wait for this batch to complete
if wait_for_completion and task_id:
console.print(
f"[cyan]Step 3: Waiting for batch {batch_start}-{batch_end} to complete...[/cyan]"
)
await wait_for_tasks_completion(client, base_url, [task_id], poll_interval)
console.print(f"[green]Batch {batch_start}-{batch_end} completed![/green]")
else:
console.print(f"[green]Batch {batch_start}-{batch_end} submitted successfully![/green]")
else:
total_processing_failures.extend([item[0] for item in successful_uploads])
console.print(f"[red]Batch {batch_start}-{batch_end} processing failed[/red]")
# Record batch failures
total_upload_failures.extend(batch_upload_failures)
# Update processed files record after each batch
if all_successful_hashes:
all_processed_files = processed_files | all_successful_hashes
save_processed_files(record_file, all_processed_files)
# Small delay between batches
if batch_end < len(files_to_upload):
console.print("[dim]Waiting 2 seconds before next batch...[/dim]")
await asyncio.sleep(2)
# Final summary
console.print("\n[bold green]=== All Batches Complete ===[/bold green]")
console.print(f" - [green]Files successfully processed:[/green] {len(total_processed_files)}")
console.print(f" - [red]Upload failures:[/red] {len(total_upload_failures)}")
if total_upload_failures:
for f in total_upload_failures:
console.print(f" - {f}")
console.print(f" - [yellow]Processing failures:[/yellow] {len(total_processing_failures)}")
if total_processing_failures:
for f in total_processing_failures:
console.print(f" - {f}")
asyncio.run(run())
"""
# Example for upload
uv run scripts/batch_upload.py upload \
--db-id your_kb_id \
--directory path/to/your/data \
--pattern "*.docx" --pattern "*.pdf" --pattern "*.html" \
--base-url http://127.0.0.1:5050/api \
--username your_username \
--password your_password \
--batch-size 20 \
--wait-for-completion \
--poll-interval 5 \
--recursive \
--record-file scripts/tmp/batch_processed_files.txt
"""
if __name__ == "__main__":
app()

File diff suppressed because it is too large Load Diff

View File

@ -1,708 +0,0 @@
"""
SQLite PostgreSQL 业务数据迁移脚本
将用户部门对话等业务数据从 SQLite 迁移到 PostgreSQL
迁移顺序按外键依赖
1. departments (无依赖)
2. users (依赖 departments)
3. conversations (依赖 users)
4. messages (依赖 conversations)
5. tool_calls (依赖 messages)
6. conversation_stats (依赖 conversations)
7. operation_logs (依赖 users)
8. message_feedbacks (依赖 messages)
9. mcp_servers (无依赖)
用法
python scripts/migrate_business_from_sqlite.py --dry-run # 预览迁移
python scripts/migrate_business_from_sqlite.py --execute # 执行迁移
python scripts/migrate_business_from_sqlite.py --verify # 验证数据
python scripts/migrate_business_from_sqlite.py --rollback # 回滚迁移
"""
import argparse
import asyncio
import os
import sys
from datetime import datetime, UTC
from typing import Any
sys.path.insert(0, os.path.dirname(os.path.dirname(__file__)))
os.environ.setdefault("YUXI_SKIP_APP_INIT", "1")
from sqlalchemy import Column, DateTime, Integer, String, Text, create_engine, select, text
from sqlalchemy.orm import declarative_base, sessionmaker
from yuxi import config
from yuxi.storage.postgres.manager import pg_manager
from yuxi.storage.postgres.models_business import (
Department,
User,
Conversation,
Message,
ToolCall,
ConversationStats,
OperationLog,
MessageFeedback,
MCPServer,
)
from yuxi.utils import logger
# ============================================================
# SQLite 模型定义 (仅用于迁移脚本,内部使用)
# ============================================================
Base = declarative_base()
class SqliteDepartment(Base):
__tablename__ = "departments"
id = Column(Integer, primary_key=True)
name = Column(String(100), nullable=False)
description = Column(Text)
created_at = Column(DateTime)
class SqliteUser(Base):
__tablename__ = "users"
id = Column(Integer, primary_key=True)
username = Column(String(50), unique=True, nullable=False)
user_id = Column(String(50), unique=True)
phone_number = Column(String(20))
avatar = Column(String(500))
password_hash = Column(String(255))
role = Column(String(20), default="user")
department_id = Column(Integer)
created_at = Column(DateTime)
last_login = Column(DateTime)
login_failed_count = Column(Integer, default=0)
last_failed_login = Column(DateTime)
login_locked_until = Column(DateTime)
is_deleted = Column(Integer, default=0) # 0=否, 1=是
deleted_at = Column(DateTime)
class SqliteConversation(Base):
__tablename__ = "conversations"
id = Column(Integer, primary_key=True)
thread_id = Column(String(50), unique=True)
user_id = Column(String(64), nullable=False)
agent_id = Column(String(50))
title = Column(String(255))
status = Column(String(20), default="active")
created_at = Column(DateTime)
updated_at = Column(DateTime)
extra_metadata = Column(Text)
class SqliteMessage(Base):
__tablename__ = "messages"
id = Column(Integer, primary_key=True)
conversation_id = Column(Integer, nullable=False)
role = Column(String(20), nullable=False)
content = Column(Text)
message_type = Column(String(20), default="text")
created_at = Column(DateTime)
token_count = Column(Integer)
extra_metadata = Column(Text)
image_content = Column(Text)
class SqliteToolCall(Base):
__tablename__ = "tool_calls"
id = Column(Integer, primary_key=True)
message_id = Column(Integer, nullable=False)
langgraph_tool_call_id = Column(String(100))
tool_name = Column(String(100))
tool_input = Column(Text)
tool_output = Column(Text)
status = Column(String(20), default="pending")
error_message = Column(Text)
created_at = Column(DateTime)
class SqliteConversationStats(Base):
__tablename__ = "conversation_stats"
id = Column(Integer, primary_key=True)
conversation_id = Column(Integer, nullable=False)
message_count = Column(Integer, default=0)
total_tokens = Column(Integer, default=0)
model_used = Column(String(100))
user_feedback = Column(String(20))
created_at = Column(DateTime)
updated_at = Column(DateTime)
class SqliteOperationLog(Base):
__tablename__ = "operation_logs"
id = Column(Integer, primary_key=True)
user_id = Column(Integer) # 外键到 users.id
operation = Column(String(100))
details = Column(Text)
ip_address = Column(String(50))
timestamp = Column(DateTime)
class SqliteMessageFeedback(Base):
__tablename__ = "message_feedbacks"
id = Column(Integer, primary_key=True)
message_id = Column(Integer, nullable=False)
user_id = Column(String(64), nullable=False)
rating = Column(String(20))
reason = Column(Text)
created_at = Column(DateTime)
class SqliteMCPServer(Base):
__tablename__ = "mcp_servers"
id = Column(Integer, primary_key=True)
name = Column(String(100), unique=True, nullable=False)
description = Column(Text)
transport = Column(String(20), default="sse")
url = Column(String(500))
command = Column(String(255))
args = Column(Text)
headers = Column(Text)
timeout = Column(Integer)
sse_read_timeout = Column(Integer)
tags = Column(Text)
icon = Column(String(500))
enabled = Column(Integer, default=1) # 1=是, 0=否
disabled_tools = Column(Text)
created_by = Column(String(100), nullable=False) # 创建人用户名
updated_by = Column(String(100), nullable=False) # 修改人用户名
created_at = Column(DateTime)
updated_at = Column(DateTime)
def _utc_dt(value: Any) -> datetime | None:
"""Convert various datetime formats to naive UTC datetime."""
if not value:
return None
if isinstance(value, datetime):
if value.tzinfo is None:
return value
return value.astimezone(UTC).replace(tzinfo=None)
if isinstance(value, (int, float)):
return datetime.fromtimestamp(value, tz=UTC).replace(tzinfo=None)
if isinstance(value, str):
v = value.strip()
if not v:
return None
try:
dt_val = datetime.fromisoformat(v.replace("Z", "+00:00"))
if dt_val.tzinfo is None:
return dt_val
return dt_val.astimezone(UTC).replace(tzinfo=None)
except ValueError:
return None
return None
class SQLiteReader:
"""SQLite 数据读取器"""
def __init__(self):
db_path = os.path.join(config.save_dir, "database", "server.db")
self.engine = create_engine(f"sqlite:///{db_path}")
self.Session = sessionmaker(bind=self.engine)
def get_session(self):
return self.Session()
def read_departments(self) -> list[SqliteDepartment]:
with self.get_session() as session:
return session.execute(select(SqliteDepartment)).scalars().all()
def read_users(self) -> list[SqliteUser]:
with self.get_session() as session:
return session.execute(select(SqliteUser)).scalars().all()
def read_conversations(self) -> list[SqliteConversation]:
with self.get_session() as session:
return session.execute(select(SqliteConversation)).scalars().all()
def read_messages(self) -> list[SqliteMessage]:
with self.get_session() as session:
return session.execute(select(SqliteMessage)).scalars().all()
def read_tool_calls(self) -> list[SqliteToolCall]:
with self.get_session() as session:
return session.execute(select(SqliteToolCall)).scalars().all()
def read_conversation_stats(self) -> list[SqliteConversationStats]:
with self.get_session() as session:
return session.execute(select(SqliteConversationStats)).scalars().all()
def read_operation_logs(self) -> list[SqliteOperationLog]:
with self.get_session() as session:
return session.execute(select(SqliteOperationLog)).scalars().all()
def read_message_feedbacks(self) -> list[SqliteMessageFeedback]:
with self.get_session() as session:
return session.execute(select(SqliteMessageFeedback)).scalars().all()
def read_mcp_servers(self) -> list[SqliteMCPServer]:
with self.get_session() as session:
return session.execute(select(SqliteMCPServer)).scalars().all()
def count_table(self, table_name: str) -> int:
with self.get_session() as session:
result = session.execute(text(f"SELECT COUNT(*) FROM {table_name}"))
return result.scalar() or 0
async def migrate_departments(sqlite_reader: SQLiteReader, dry_run: bool, execute: bool) -> dict[str, int]:
"""迁移部门数据"""
sqlite_depts = sqlite_reader.read_departments()
logger.info(f"准备迁移 {len(sqlite_depts)} 个部门")
created = 0
if dry_run:
for sqlite_dept in sqlite_depts:
logger.info(f"[DRY-RUN] 将创建部门: {sqlite_dept.name}")
elif execute:
async with pg_manager.get_async_session_context() as session:
for sqlite_dept in sqlite_depts:
# 检查是否已存在
existing = await session.execute(select(Department).where(Department.id == sqlite_dept.id))
if existing.scalar_one_or_none() is None:
dept = Department(
id=sqlite_dept.id,
name=sqlite_dept.name,
description=sqlite_dept.description,
created_at=_utc_dt(sqlite_dept.created_at),
)
session.add(dept)
created += 1
return {"total": len(sqlite_depts), "created": created}
async def migrate_users(sqlite_reader: SQLiteReader, dry_run: bool, execute: bool) -> dict[str, int]:
"""迁移用户数据"""
sqlite_users = sqlite_reader.read_users()
logger.info(f"准备迁移 {len(sqlite_users)} 个用户")
created = 0
if dry_run:
for sqlite_user in sqlite_users:
logger.info(f"[DRY-RUN] 将创建用户: {sqlite_user.username} ({sqlite_user.user_id})")
elif execute:
async with pg_manager.get_async_session_context() as session:
for sqlite_user in sqlite_users:
existing = await session.execute(select(User).where(User.id == sqlite_user.id))
if existing.scalar_one_or_none() is None:
user = User(
id=sqlite_user.id,
username=sqlite_user.username,
user_id=sqlite_user.user_id,
phone_number=sqlite_user.phone_number,
avatar=sqlite_user.avatar,
password_hash=sqlite_user.password_hash,
role=sqlite_user.role,
department_id=sqlite_user.department_id,
created_at=_utc_dt(sqlite_user.created_at),
last_login=_utc_dt(sqlite_user.last_login),
login_failed_count=sqlite_user.login_failed_count,
last_failed_login=_utc_dt(sqlite_user.last_failed_login),
login_locked_until=_utc_dt(sqlite_user.login_locked_until),
is_deleted=sqlite_user.is_deleted,
deleted_at=_utc_dt(sqlite_user.deleted_at),
)
session.add(user)
created += 1
return {"total": len(sqlite_users), "created": created}
async def migrate_conversations(sqlite_reader: SQLiteReader, dry_run: bool, execute: bool) -> dict[str, int]:
"""迁移对话数据"""
sqlite_convs = sqlite_reader.read_conversations()
logger.info(f"准备迁移 {len(sqlite_convs)} 个对话")
created = 0
if dry_run:
for sqlite_conv in sqlite_convs:
logger.info(f"[DRY-RUN] 将创建对话: {sqlite_conv.thread_id}")
elif execute:
async with pg_manager.get_async_session_context() as session:
for sqlite_conv in sqlite_convs:
existing = await session.execute(select(Conversation).where(Conversation.id == sqlite_conv.id))
if existing.scalar_one_or_none() is None:
# 截断过长的 title
title = sqlite_conv.title
if title and len(title) > 255:
title = title[:255]
logger.warning(f"截断对话标题 (id={sqlite_conv.id}): 原始长度={len(sqlite_conv.title)}")
conv = Conversation(
id=sqlite_conv.id,
thread_id=sqlite_conv.thread_id,
user_id=sqlite_conv.user_id,
agent_id=sqlite_conv.agent_id,
title=title,
status=sqlite_conv.status,
created_at=_utc_dt(sqlite_conv.created_at),
updated_at=_utc_dt(sqlite_conv.updated_at),
extra_metadata=sqlite_conv.extra_metadata,
)
session.add(conv)
created += 1
return {"total": len(sqlite_convs), "created": created}
async def migrate_messages(sqlite_reader: SQLiteReader, dry_run: bool, execute: bool) -> dict[str, int]:
"""迁移消息数据"""
sqlite_messages = sqlite_reader.read_messages()
logger.info(f"准备迁移 {len(sqlite_messages)} 条消息")
created = 0
if dry_run:
for sqlite_msg in sqlite_messages:
logger.info(f"[DRY-RUN] 将创建消息: id={sqlite_msg.id}, conversation={sqlite_msg.conversation_id}")
elif execute:
async with pg_manager.get_async_session_context() as session:
for sqlite_msg in sqlite_messages:
existing = await session.execute(select(Message).where(Message.id == sqlite_msg.id))
if existing.scalar_one_or_none() is None:
msg = Message(
id=sqlite_msg.id,
conversation_id=sqlite_msg.conversation_id,
role=sqlite_msg.role,
content=sqlite_msg.content,
message_type=sqlite_msg.message_type,
created_at=_utc_dt(sqlite_msg.created_at),
token_count=sqlite_msg.token_count,
extra_metadata=sqlite_msg.extra_metadata,
image_content=sqlite_msg.image_content,
)
session.add(msg)
created += 1
return {"total": len(sqlite_messages), "created": created}
async def migrate_tool_calls(sqlite_reader: SQLiteReader, dry_run: bool, execute: bool) -> dict[str, int]:
"""迁移工具调用数据"""
sqlite_calls = sqlite_reader.read_tool_calls()
logger.info(f"准备迁移 {len(sqlite_calls)} 个工具调用")
created = 0
if dry_run:
for sqlite_call in sqlite_calls:
logger.info(f"[DRY-RUN] 将创建工具调用: id={sqlite_call.id}, tool={sqlite_call.tool_name}")
elif execute:
async with pg_manager.get_async_session_context() as session:
for sqlite_call in sqlite_calls:
existing = await session.execute(select(ToolCall).where(ToolCall.id == sqlite_call.id))
if existing.scalar_one_or_none() is None:
call = ToolCall(
id=sqlite_call.id,
message_id=sqlite_call.message_id,
langgraph_tool_call_id=sqlite_call.langgraph_tool_call_id,
tool_name=sqlite_call.tool_name,
tool_input=sqlite_call.tool_input,
tool_output=sqlite_call.tool_output,
status=sqlite_call.status,
error_message=sqlite_call.error_message,
created_at=_utc_dt(sqlite_call.created_at),
)
session.add(call)
created += 1
return {"total": len(sqlite_calls), "created": created}
async def migrate_conversation_stats(sqlite_reader: SQLiteReader, dry_run: bool, execute: bool) -> dict[str, int]:
"""迁移对话统计数据"""
sqlite_stats = sqlite_reader.read_conversation_stats()
logger.info(f"准备迁移 {len(sqlite_stats)} 条对话统计")
created = 0
if dry_run:
for sqlite_stat in sqlite_stats:
logger.info(f"[DRY-RUN] 将创建对话统计: conversation_id={sqlite_stat.conversation_id}")
elif execute:
async with pg_manager.get_async_session_context() as session:
for sqlite_stat in sqlite_stats:
existing = await session.execute(
select(ConversationStats).where(ConversationStats.id == sqlite_stat.id)
)
if existing.scalar_one_or_none() is None:
stat = ConversationStats(
id=sqlite_stat.id,
conversation_id=sqlite_stat.conversation_id,
message_count=sqlite_stat.message_count,
total_tokens=sqlite_stat.total_tokens,
model_used=sqlite_stat.model_used,
user_feedback=sqlite_stat.user_feedback,
created_at=_utc_dt(sqlite_stat.created_at),
updated_at=_utc_dt(sqlite_stat.updated_at),
)
session.add(stat)
created += 1
return {"total": len(sqlite_stats), "created": created}
async def migrate_operation_logs(sqlite_reader: SQLiteReader, dry_run: bool, execute: bool) -> dict[str, int]:
"""迁移操作日志数据"""
sqlite_logs = sqlite_reader.read_operation_logs()
logger.info(f"准备迁移 {len(sqlite_logs)} 条操作日志")
created = 0
if dry_run:
for sqlite_log in sqlite_logs:
logger.info(f"[DRY-RUN] 将创建操作日志: id={sqlite_log.id}, operation={sqlite_log.operation}")
elif execute:
async with pg_manager.get_async_session_context() as session:
for sqlite_log in sqlite_logs:
existing = await session.execute(select(OperationLog).where(OperationLog.id == sqlite_log.id))
if existing.scalar_one_or_none() is None:
log = OperationLog(
id=sqlite_log.id,
user_id=sqlite_log.user_id,
operation=sqlite_log.operation,
details=sqlite_log.details,
ip_address=sqlite_log.ip_address,
timestamp=_utc_dt(sqlite_log.timestamp),
)
session.add(log)
created += 1
return {"total": len(sqlite_logs), "created": created}
async def migrate_message_feedbacks(sqlite_reader: SQLiteReader, dry_run: bool, execute: bool) -> dict[str, int]:
"""迁移消息反馈数据"""
sqlite_feedbacks = sqlite_reader.read_message_feedbacks()
logger.info(f"准备迁移 {len(sqlite_feedbacks)} 条消息反馈")
created = 0
if dry_run:
for sqlite_fb in sqlite_feedbacks:
logger.info(f"[DRY-RUN] 将创建消息反馈: id={sqlite_fb.id}, rating={sqlite_fb.rating}")
elif execute:
async with pg_manager.get_async_session_context() as session:
for sqlite_fb in sqlite_feedbacks:
existing = await session.execute(select(MessageFeedback).where(MessageFeedback.id == sqlite_fb.id))
if existing.scalar_one_or_none() is None:
fb = MessageFeedback(
id=sqlite_fb.id,
message_id=sqlite_fb.message_id,
user_id=sqlite_fb.user_id,
rating=sqlite_fb.rating,
reason=sqlite_fb.reason,
created_at=_utc_dt(sqlite_fb.created_at),
)
session.add(fb)
created += 1
return {"total": len(sqlite_feedbacks), "created": created}
async def migrate_mcp_servers(sqlite_reader: SQLiteReader, dry_run: bool, execute: bool) -> dict[str, int]:
"""迁移 MCP 服务器数据"""
sqlite_servers = sqlite_reader.read_mcp_servers()
logger.info(f"准备迁移 {len(sqlite_servers)} 个 MCP 服务器")
created = 0
if dry_run:
for sqlite_server in sqlite_servers:
logger.info(f"[DRY-RUN] 将创建 MCP 服务器: {sqlite_server.name}")
elif execute:
async with pg_manager.get_async_session_context() as session:
for sqlite_server in sqlite_servers:
existing = await session.execute(select(MCPServer).where(MCPServer.name == sqlite_server.name))
if existing.scalar_one_or_none() is None:
server = MCPServer(
name=sqlite_server.name,
description=sqlite_server.description,
transport=sqlite_server.transport,
url=sqlite_server.url,
command=sqlite_server.command,
args=sqlite_server.args,
env=getattr(sqlite_server, "env", None),
headers=sqlite_server.headers,
timeout=sqlite_server.timeout,
sse_read_timeout=sqlite_server.sse_read_timeout,
tags=sqlite_server.tags,
icon=sqlite_server.icon,
enabled=sqlite_server.enabled,
disabled_tools=sqlite_server.disabled_tools,
created_by=sqlite_server.created_by,
updated_by=sqlite_server.updated_by,
created_at=_utc_dt(sqlite_server.created_at),
updated_at=_utc_dt(sqlite_server.updated_at),
)
session.add(server)
created += 1
return {"total": len(sqlite_servers), "created": created}
async def verify_migration(sqlite_reader: SQLiteReader) -> dict[str, dict]:
"""验证迁移结果"""
# 使用 (模型, 主键列名) 格式,支持不同表使用不同的主键
tables = [
("departments", Department, "id"),
("users", User, "id"),
("conversations", Conversation, "id"),
("messages", Message, "id"),
("tool_calls", ToolCall, "id"),
("conversation_stats", ConversationStats, "id"),
("operation_logs", OperationLog, "id"),
("message_feedbacks", MessageFeedback, "id"),
("mcp_servers", MCPServer, "name"), # MCPServer 使用 name 作为主键
]
results = {}
for table_name, model, pk_column in tables:
sqlite_count = sqlite_reader.count_table(table_name)
async with pg_manager.get_async_session_context() as session:
from sqlalchemy import func
pk_attr = getattr(model, pk_column)
result = await session.execute(select(func.count(pk_attr)))
pg_count = result.scalar() or 0
results[table_name] = {
"sqlite": sqlite_count,
"postgresql": pg_count,
"match": sqlite_count == pg_count,
}
return results
async def rollback_migration() -> None:
"""回滚迁移 - 删除所有业务数据表"""
logger.warning("开始回滚迁移...")
# 按外键依赖顺序删除
tables_to_delete = [
MessageFeedback,
OperationLog,
ConversationStats,
ToolCall,
Message,
Conversation,
User,
Department,
MCPServer,
]
for model in tables_to_delete:
async with pg_manager.get_async_session_context() as session:
result = await session.execute(select(model))
records = result.scalars().all()
for record in records:
await session.delete(record)
logger.warning("回滚完成 - 已删除所有迁移的业务数据")
async def migrate_all(sqlite_reader: SQLiteReader, dry_run: bool, execute: bool) -> dict[str, Any]:
"""执行所有迁移"""
results = {}
# 按外键依赖顺序迁移
results["departments"] = await migrate_departments(sqlite_reader, dry_run, execute)
results["users"] = await migrate_users(sqlite_reader, dry_run, execute)
results["conversations"] = await migrate_conversations(sqlite_reader, dry_run, execute)
results["messages"] = await migrate_messages(sqlite_reader, dry_run, execute)
results["tool_calls"] = await migrate_tool_calls(sqlite_reader, dry_run, execute)
results["conversation_stats"] = await migrate_conversation_stats(sqlite_reader, dry_run, execute)
results["operation_logs"] = await migrate_operation_logs(sqlite_reader, dry_run, execute)
results["message_feedbacks"] = await migrate_message_feedbacks(sqlite_reader, dry_run, execute)
results["mcp_servers"] = await migrate_mcp_servers(sqlite_reader, dry_run, execute)
return results
async def main() -> None:
parser = argparse.ArgumentParser(description="SQLite 到 PostgreSQL 业务数据迁移")
parser.add_argument("--dry-run", action="store_true", help="预览迁移,不执行")
parser.add_argument("--execute", action="store_true", help="执行迁移")
parser.add_argument("--verify", action="store_true", help="验证迁移结果")
parser.add_argument("--rollback", action="store_true", help="回滚迁移")
parser.add_argument("--migrate-all", action="store_true", help="迁移所有业务数据")
parser.add_argument("--init-tables", action="store_true", help="仅初始化业务表结构")
args = parser.parse_args()
if not any([args.dry_run, args.execute, args.verify, args.rollback, args.migrate_all, args.init_tables]):
args.dry_run = True
# 初始化 PostgreSQL 管理器
pg_manager.initialize()
logger.info("PostgreSQL manager initialized")
if args.init_tables:
# 仅初始化表结构
await pg_manager.create_business_tables()
logger.info("业务表结构初始化完成")
return
if args.verify:
# 验证模式
sqlite_reader = SQLiteReader()
results = await verify_migration(sqlite_reader)
logger.info("=" * 60)
logger.info("迁移验证结果:")
logger.info("=" * 60)
all_match = True
for table_name, counts in results.items():
status = "" if counts["match"] else ""
logger.info(f"{status} {table_name}: SQLite={counts['sqlite']}, PostgreSQL={counts['postgresql']}")
if not counts["match"]:
all_match = False
logger.info("=" * 60)
logger.info(f"全部匹配: {'' if all_match else ''}")
return
if args.rollback:
# 回滚模式
if args.dry_run:
logger.info("[DRY-RUN] 将回滚所有迁移的业务数据")
else:
await rollback_migration()
return
# 迁移模式
sqlite_reader = SQLiteReader()
if args.migrate_all:
# 检查是否需要初始化表结构
logger.info("检查业务表结构...")
await pg_manager.create_business_tables()
logger.info("业务表结构就绪")
results = await migrate_all(sqlite_reader, args.dry_run, args.execute)
logger.info("=" * 60)
logger.info("迁移完成:")
for table_name, counts in results.items():
logger.info(f" {table_name}: {counts['created']}/{counts['total']}")
logger.info("=" * 60)
if not args.dry_run:
logger.info("建议运行 --verify 验证数据完整性")
else:
logger.info("使用 --migrate-all 执行迁移,或使用 --verify 验证数据")
if __name__ == "__main__":
asyncio.run(main())

View File

@ -1,339 +0,0 @@
import argparse
import asyncio
import glob
import json
import os
import sys
from datetime import datetime, UTC
from typing import Any
sys.path.insert(0, os.path.dirname(os.path.dirname(__file__)))
os.environ.setdefault("YUXI_SKIP_APP_INIT", "1")
from yuxi import config
from yuxi.repositories.evaluation_repository import EvaluationRepository
from yuxi.repositories.knowledge_base_repository import KnowledgeBaseRepository
from yuxi.repositories.knowledge_file_repository import KnowledgeFileRepository
from yuxi.repositories.task_repository import TaskRepository
from yuxi.utils import logger
def _load_json(path: str) -> dict[str, Any]:
if not os.path.exists(path):
return {}
with open(path, encoding="utf-8") as f:
return json.load(f)
def _utc_dt(value: Any) -> datetime | None:
"""Convert various datetime formats to naive UTC datetime (consistent with model)."""
if not value:
return None
if isinstance(value, datetime):
# 转换为 UTC 并移除时区信息(模型使用 DateTime 无时区)
if value.tzinfo is None:
return value
return value.astimezone(UTC).replace(tzinfo=None)
if isinstance(value, (int, float)):
# 时间戳转换为 UTC 时间
return datetime.fromtimestamp(value, tz=UTC).replace(tzinfo=None)
if isinstance(value, str):
v = value.strip()
if not v:
return None
try:
# 解析 ISO 格式并转换为 UTC
dt_val = datetime.fromisoformat(v.replace("Z", "+00:00"))
if dt_val.tzinfo is None:
return dt_val
return dt_val.astimezone(UTC).replace(tzinfo=None)
except ValueError:
return None
return None
def _default_share_config(meta: dict[str, Any]) -> dict[str, Any]:
share_config = meta.get("share_config") or {}
if "is_shared" not in share_config:
share_config["is_shared"] = True
if "accessible_departments" not in share_config:
share_config["accessible_departments"] = []
return share_config
async def rollback_all() -> None:
eval_repo = EvaluationRepository()
kb_repo = KnowledgeBaseRepository()
file_repo = KnowledgeFileRepository()
task_repo = TaskRepository()
await task_repo.delete_all()
await eval_repo.delete_all()
rows = await kb_repo.get_all()
for row in rows:
await file_repo.delete_by_db_id(row.db_id)
await kb_repo.delete(row.db_id)
async def migrate(dry_run: bool, execute: bool, rollback: bool) -> None:
from yuxi.storage.postgres.manager import pg_manager
base_dir = os.path.join(config.save_dir, "knowledge_base_data")
global_meta_path = os.path.join(base_dir, "global_metadata.json")
global_meta = _load_json(global_meta_path).get("databases", {})
if rollback:
if dry_run:
logger.info("Dry-run rollback: would delete all knowledge metadata tables")
return
await rollback_all()
logger.info("Rollback completed")
return
# 初始化表结构
pg_manager.initialize()
await pg_manager.create_tables()
logger.info("知识库表结构初始化完成")
kb_repo = KnowledgeBaseRepository()
file_repo = KnowledgeFileRepository()
eval_repo = EvaluationRepository()
task_repo = TaskRepository()
kb_rows: list[dict[str, Any]] = []
file_rows: list[tuple[str, dict[str, Any]]] = []
benchmark_rows: list[dict[str, Any]] = []
result_rows: list[dict[str, Any]] = []
result_detail_rows: list[tuple[str, int, dict[str, Any]]] = []
kb_type_dirs = [
p for p in glob.glob(os.path.join(base_dir, "*_data")) if os.path.isdir(p) and os.path.basename(p) != "uploads"
]
for kb_dir in kb_type_dirs:
kb_type = os.path.basename(kb_dir)[: -len("_data")]
meta_file = os.path.join(kb_dir, f"metadata_{kb_type}.json")
meta = _load_json(meta_file)
databases_meta: dict[str, Any] = meta.get("databases", {})
files_meta: dict[str, Any] = meta.get("files", {})
benchmarks_meta: dict[str, Any] = meta.get("benchmarks", {})
for db_id, db_meta in databases_meta.items():
g = global_meta.get(db_id, {})
created_at = _utc_dt(g.get("created_at") or db_meta.get("created_at"))
updated_at = _utc_dt(g.get("updated_at")) or created_at
kb_rows.append(
{
"db_id": db_id,
"name": g.get("name") or db_meta.get("name") or db_id,
"description": g.get("description") or db_meta.get("description"),
"kb_type": g.get("kb_type") or db_meta.get("kb_type") or kb_type,
"embed_info": db_meta.get("embed_info") or g.get("embed_info"),
"llm_info": db_meta.get("llm_info") or g.get("llm_info"),
"query_params": db_meta.get("query_params") or g.get("query_params"),
"additional_params": g.get("additional_params") or db_meta.get("metadata") or {},
"share_config": _default_share_config(g or {}),
"mindmap": g.get("mindmap"),
"sample_questions": g.get("sample_questions") or [],
"created_at": created_at,
"updated_at": updated_at,
}
)
for file_id, fmeta in files_meta.items():
db_id = fmeta.get("database_id")
if not db_id:
continue
file_rows.append(
(
file_id,
{
"db_id": db_id,
"parent_id": fmeta.get("parent_id"),
"filename": fmeta.get("filename") or "",
"original_filename": fmeta.get("original_filename") or fmeta.get("file_name"),
"file_type": fmeta.get("file_type") or fmeta.get("type"),
"path": fmeta.get("path"),
"minio_url": fmeta.get("minio_url"),
"markdown_file": fmeta.get("markdown_file"),
"status": fmeta.get("status"),
"content_hash": fmeta.get("content_hash"),
"file_size": fmeta.get("size") or fmeta.get("file_size"),
"content_type": fmeta.get("content_type"),
"processing_params": fmeta.get("processing_params"),
"is_folder": bool(fmeta.get("is_folder", False)),
"error_message": fmeta.get("error") or fmeta.get("error_message"),
"created_by": str(fmeta.get("created_by")) if fmeta.get("created_by") else None,
"updated_by": str(fmeta.get("updated_by")) if fmeta.get("updated_by") else None,
"created_at": _utc_dt(fmeta.get("created_at")),
"updated_at": _utc_dt(fmeta.get("updated_at")) or _utc_dt(fmeta.get("created_at")),
},
)
)
for db_id, bmap in benchmarks_meta.items():
if not isinstance(bmap, dict):
continue
for benchmark_id, bmeta in bmap.items():
benchmark_rows.append(
{
"benchmark_id": benchmark_id,
"db_id": db_id,
"name": bmeta.get("name") or benchmark_id,
"description": bmeta.get("description"),
"question_count": int(bmeta.get("question_count") or 0),
"has_gold_chunks": bool(bmeta.get("has_gold_chunks")),
"has_gold_answers": bool(bmeta.get("has_gold_answers")),
"data_file_path": bmeta.get("benchmark_file") or bmeta.get("data_file_path"),
"created_by": str(bmeta.get("created_by")) if bmeta.get("created_by") else None,
"created_at": _utc_dt(bmeta.get("created_at")),
"updated_at": _utc_dt(bmeta.get("updated_at")) or _utc_dt(bmeta.get("created_at")),
}
)
for db_id in databases_meta.keys():
result_dir = os.path.join(kb_dir, db_id, "results")
if not os.path.isdir(result_dir):
continue
for result_path in glob.glob(os.path.join(result_dir, "*.json")):
try:
data = _load_json(result_path)
except Exception as exc:
logger.warning(f"Skip invalid result file {result_path}: {exc}")
continue
task_id = data.get("task_id") or os.path.splitext(os.path.basename(result_path))[0]
benchmark_id = data.get("benchmark_id")
started_at = _utc_dt(data.get("started_at"))
result_rows.append(
{
"task_id": task_id,
"db_id": db_id,
"benchmark_id": benchmark_id,
"status": data.get("status") or "completed",
"retrieval_config": data.get("retrieval_config") or {},
"metrics": data.get("metrics") or {},
"overall_score": data.get("overall_score"),
"total_questions": int(data.get("total_questions") or 0),
"completed_questions": int(data.get("completed_questions") or 0),
"started_at": started_at,
"completed_at": _utc_dt(data.get("completed_at")) or started_at,
"created_by": str(data.get("created_by")) if data.get("created_by") else None,
}
)
interim = data.get("interim_results") or data.get("results") or []
for idx, item in enumerate(interim):
result_detail_rows.append(
(
task_id,
idx,
{
"query_text": item.get("query") or item.get("query_text") or "",
"gold_chunk_ids": item.get("gold_chunk_ids"),
"gold_answer": item.get("gold_answer"),
"generated_answer": item.get("generated_answer"),
"retrieved_chunks": item.get("retrieved_chunks"),
"metrics": item.get("metrics") or {},
},
)
)
tasks_json_path = os.path.join(config.save_dir, "tasks", "tasks.json")
task_rows: list[dict[str, Any]] = _load_json(tasks_json_path).get("tasks", []) or []
logger.info(
f"Prepared: knowledge_bases={len(kb_rows)}, knowledge_files={len(file_rows)}, "
f"benchmarks={len(benchmark_rows)}, results={len(result_rows)}, result_details={len(result_detail_rows)}, "
f"tasks={len(task_rows)}"
)
if dry_run and not execute:
return
for payload in kb_rows:
db_id = payload["db_id"]
existing = await kb_repo.get_by_id(db_id)
data = payload.copy()
if existing is None:
await kb_repo.create(data)
else:
await kb_repo.update(db_id, data)
# 先插入文件夹,再插入普通文件(确保父文件夹先存在)
folders = [(fid, data) for fid, data in file_rows if data.get("is_folder")]
files = [(fid, data) for fid, data in file_rows if not data.get("is_folder")]
for file_id, data in folders:
await file_repo.upsert(file_id=file_id, data=data)
for file_id, data in files:
await file_repo.upsert(file_id=file_id, data=data)
for payload in benchmark_rows:
# 检查知识库是否存在
kb = await kb_repo.get_by_id(payload["db_id"])
if kb is None:
logger.warning(f"Skipping benchmark {payload['benchmark_id']}: knowledge base {payload['db_id']} not found")
continue
existing = await eval_repo.get_benchmark(payload["benchmark_id"])
if existing is None:
await eval_repo.create_benchmark(payload)
for payload in result_rows:
# 检查知识库是否存在
kb = await kb_repo.get_by_id(payload["db_id"])
if kb is None:
logger.warning(f"Skipping result {payload['task_id']}: knowledge base {payload['db_id']} not found")
continue
existing = await eval_repo.get_result(payload["task_id"])
if existing is None:
await eval_repo.create_result(payload)
else:
await eval_repo.update_result(payload["task_id"], payload)
for task_id, idx, data in result_detail_rows:
await eval_repo.upsert_result_detail(task_id=task_id, query_index=idx, data=data)
for item in task_rows:
task_id = item.get("id")
if not task_id:
continue
payload = item.get("payload") or {}
result = item.get("result")
await task_repo.upsert(
task_id,
{
"name": item.get("name") or "Unnamed Task",
"type": item.get("type") or "general",
"status": item.get("status") or "pending",
"progress": float(item.get("progress") or 0.0),
"message": item.get("message") or "",
"payload": payload,
"result": result,
"error": item.get("error"),
"cancel_requested": 1 if item.get("cancel_requested") else 0,
"created_at": _utc_dt(item.get("created_at")),
"updated_at": _utc_dt(item.get("updated_at")) or _utc_dt(item.get("created_at")),
"started_at": _utc_dt(item.get("started_at")),
"completed_at": _utc_dt(item.get("completed_at")),
},
)
logger.info("Migration completed")
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--dry-run", action="store_true")
parser.add_argument("--execute", action="store_true")
parser.add_argument("--rollback", action="store_true")
args = parser.parse_args()
if not args.dry_run and not args.execute and not args.rollback:
args.dry_run = True
asyncio.run(migrate(dry_run=args.dry_run, execute=args.execute, rollback=args.rollback))
if __name__ == "__main__":
main()

View File

@ -1,81 +0,0 @@
import json
import random
import re
from pathlib import Path
import pandas as pd
import typer
app = typer.Typer()
def sanitize_filename(name: str) -> str:
return re.sub(r'[\\/*?:"<>|]', "_", str(name).strip())
def random_suffix() -> str:
return f"_{random.randint(10000000, 99999999)}"
def read_table(file_path: Path) -> pd.DataFrame:
suffix = file_path.suffix.lower()
if suffix in [".xlsx", ".xls"]:
return pd.read_excel(file_path)
elif suffix == ".csv":
return pd.read_csv(file_path)
elif suffix == ".json":
with open(file_path, encoding="utf-8") as f:
data = json.load(f)
if isinstance(data, list) and len(data) > 1 and isinstance(data[0], dict):
return pd.DataFrame(data)
else:
raise ValueError("JSON 文件格式不符合要求:应为元素个数 > 1 的数组,每个元素是对象。")
else:
raise ValueError(f"不支持的文件格式:{suffix}")
def export_txts(df: pd.DataFrame, output_dir: Path, title_field: str = "标题"):
output_dir.mkdir(parents=True, exist_ok=True)
df.columns = [c.strip() for c in df.columns]
if title_field not in df.columns:
title_field = df.columns[0] # fallback
for idx, row in df.iterrows():
title = str(row.get(title_field, "")).strip()
if not title:
title = f"{str(row[df.columns[0]])}{random_suffix()}"
else:
title = sanitize_filename(title)
filename = f"{title}.txt"
file_path = output_dir / filename
# 构造内容:字段: 值,每行一个
content = "\n".join(f"{col}: {row[col]}" for col in df.columns)
with open(file_path, "w", encoding="utf-8") as f:
f.write(content)
typer.echo(f"✅ 成功导出 {len(df)} 个文件到目录:{output_dir}")
@app.command()
def convert(
input_file: Path = typer.Argument(..., help="输入文件(.xlsx/.xls/.csv/.json"),
out_dir: Path = typer.Option("output", help="输出目录"),
title_field: str = typer.Option("标题", help="标题字段名(用于文件名)"),
):
"""
将结构化数据文件Excel/CSV/JSON转换为多个 .txt 文件
"""
try:
df = read_table(input_file)
export_txts(df, out_dir, title_field)
except Exception as e:
typer.echo(f"❌ 错误:{e}", err=True)
raise typer.Exit(code=1)
if __name__ == "__main__":
app()

View File

@ -1,82 +0,0 @@
import os
from pymilvus import Collection, connections, utility
def get_collection_info(collection_name, alias):
"""Safely gets a collection object and its number of entities."""
try:
collection = Collection(collection_name, using=alias)
collection.load()
return collection, collection.num_entities
except Exception as e:
print(f"Error getting info for collection '{collection_name}': {e}")
return None, 0
def rename_and_resolve_duplicates():
"""
Connects to Milvus, renames collections from 'kb_kb_' to 'kb_',
and resolves duplicates by keeping the collection with more rows.
"""
milvus_uri = os.getenv("MILVUS_URI") or "http://localhost:19530"
milvus_token = os.getenv("MILVUS_TOKEN") or ""
connection_alias = "rename_script"
try:
print(f"Connecting to Milvus at {milvus_uri}...")
connections.connect(alias=connection_alias, uri=milvus_uri, token=milvus_token)
print("Successfully connected to Milvus.")
all_collections = utility.list_collections(using=connection_alias)
collections_to_rename = [c for c in all_collections if c.startswith("kb_kb_")]
if not collections_to_rename:
print("No collections with the prefix 'kb_kb_' found. Nothing to do.")
return
print(f"Found {len(collections_to_rename)} collections with 'kb_kb_' prefix to process.")
for old_name in collections_to_rename:
new_name = old_name.replace("kb_kb_", "kb_", 1)
try:
print(f"Attempting to rename '{old_name}' to '{new_name}'...")
utility.rename_collection(old_name, new_name, using=connection_alias)
print(f"Successfully renamed '{old_name}' to '{new_name}'.")
except Exception as e:
# Check if it's a duplicate name error
if "duplicated new collection name" in str(e):
print(f"Rename failed: Target collection '{new_name}' already exists. Resolving duplicate...")
# Get info for both collections
old_coll, old_count = get_collection_info(old_name, connection_alias)
new_coll, new_count = get_collection_info(new_name, connection_alias)
print(f"Comparing row counts: '{old_name}' ({old_count} rows) vs '{new_name}' ({new_count} rows).")
if old_count > new_count:
print(f"'{old_name}' has more rows. Deleting '{new_name}' and retrying rename.")
utility.drop_collection(new_name, using=connection_alias)
print(f"Dropped collection '{new_name}'.")
# Retry renaming
utility.rename_collection(old_name, new_name, using=connection_alias)
print(f"Successfully renamed '{old_name}' to '{new_name}'.")
else:
print(f"'{new_name}' has more or equal rows. Deleting '{old_name}'.")
utility.drop_collection(old_name, using=connection_alias)
print(f"Dropped collection '{old_name}'.")
else:
print(f"An unexpected error occurred while renaming '{old_name}': {e}")
print("\nProcess finished.")
except Exception as e:
print(f"A critical error occurred: {e}")
finally:
if connection_alias in connections.list_connections():
connections.disconnect(connection_alias)
print("Disconnected from Milvus.")
if __name__ == "__main__":
rename_and_resolve_duplicates()

View File

@ -1,236 +0,0 @@
import os
import sys
import time
from dataclasses import dataclass
import requests
@dataclass
class Account:
username: str
password: str
label: str
ACCOUNTS = {
"superadmin": Account("zwj", "zwj12138", "superadmin"),
"dept_admin": Account("ceshizhuguan", "test_admin123", "dept_admin"),
"dept_user": Account("food2025", "jnufood", "dept_user"),
}
def _base_url() -> str:
return os.getenv("BASE_URL", "http://localhost:8000").rstrip("/")
def _request(method: str, path: str, *, token: str | None = None, json_data=None):
url = f"{_base_url()}{path}"
headers = {}
if token:
headers["Authorization"] = f"Bearer {token}"
resp = requests.request(method, url, headers=headers, json=json_data, timeout=60)
return resp
def login(account: Account) -> str:
url = f"{_base_url()}/api/auth/token"
resp = requests.post(
url,
data={"username": account.username, "password": account.password},
headers={"Content-Type": "application/x-www-form-urlencoded"},
timeout=60,
)
if resp.status_code != 200:
raise RuntimeError(f"login failed {account.label}: {resp.status_code} {resp.text}")
data = resp.json()
return data["access_token"], data.get("department_id"), data.get("user_id")
def get_first_agent_id(token: str) -> str:
resp = _request("GET", "/api/chat/agent", token=token)
if resp.status_code != 200:
raise RuntimeError(f"get agents failed: {resp.status_code} {resp.text}")
agents = resp.json().get("agents", [])
if not agents:
raise RuntimeError("no agents returned")
return agents[0]["id"]
def list_configs(token: str, agent_id: str) -> list[dict]:
resp = _request("GET", f"/api/chat/agent/{agent_id}/configs", token=token)
if resp.status_code != 200:
raise RuntimeError(f"list configs failed: {resp.status_code} {resp.text}")
return resp.json().get("configs", [])
def get_config(token: str, agent_id: str, config_id: int) -> dict:
resp = _request("GET", f"/api/chat/agent/{agent_id}/configs/{config_id}", token=token)
if resp.status_code != 200:
raise RuntimeError(f"get config failed: {resp.status_code} {resp.text}")
return resp.json()["config"]
def create_config(token: str, agent_id: str, name: str, set_default: bool = False) -> dict:
payload = {
"name": name,
"description": f"created-by-test {int(time.time())}",
"icon": None,
"pics": [],
"examples": ["hello"],
"config_json": {"context": {"system_prompt": f"system_prompt::{name}"}},
"set_default": set_default,
}
resp = _request("POST", f"/api/chat/agent/{agent_id}/configs", token=token, json_data=payload)
if resp.status_code != 200:
raise RuntimeError(f"create config failed: {resp.status_code} {resp.text}")
return resp.json()["config"]
def update_config(token: str, agent_id: str, config_id: int, context_updates: dict) -> dict:
payload = {"config_json": {"context": context_updates}}
resp = _request("PUT", f"/api/chat/agent/{agent_id}/configs/{config_id}", token=token, json_data=payload)
if resp.status_code != 200:
raise RuntimeError(f"update config failed: {resp.status_code} {resp.text}")
return resp.json()["config"]
def set_default(token: str, agent_id: str, config_id: int) -> dict:
resp = _request("POST", f"/api/chat/agent/{agent_id}/configs/{config_id}/set_default", token=token, json_data={})
if resp.status_code != 200:
raise RuntimeError(f"set default failed: {resp.status_code} {resp.text}")
return resp.json()["config"]
def delete_config(token: str, agent_id: str, config_id: int) -> None:
resp = _request("DELETE", f"/api/chat/agent/{agent_id}/configs/{config_id}", token=token)
if resp.status_code != 200:
raise RuntimeError(f"delete config failed: {resp.status_code} {resp.text}")
def chat_smoke(token: str, agent_id: str, config_id: int) -> None:
url = f"{_base_url()}/api/chat/agent/{agent_id}"
resp = requests.post(
url,
headers={"Authorization": f"Bearer {token}", "Content-Type": "application/json"},
json={"query": "ping", "config": {"thread_id": None, "agent_config_id": config_id}},
stream=True,
timeout=120,
)
if resp.status_code != 200:
raise RuntimeError(f"chat failed: {resp.status_code} {resp.text}")
lines = []
for line in resp.iter_lines(decode_unicode=True):
if not line:
continue
lines.append(line)
if '"status": "finished"' in line:
break
if not any('"status": "init"' in s for s in lines):
raise RuntimeError("chat stream missing init chunk")
if not any('"status": "finished"' in s for s in lines):
raise RuntimeError("chat stream missing finished chunk")
def assert_forbidden(resp: requests.Response, label: str):
if resp.status_code != 403:
raise RuntimeError(f"expected 403 for {label}, got {resp.status_code}: {resp.text}")
def main():
super_token, super_dept_id, super_user_id = login(ACCOUNTS["superadmin"])
test_admin_token, test_admin_dept_id, _ = login(ACCOUNTS["dept_admin"])
default_dept_token, default_dept_id, _ = login(ACCOUNTS["dept_user"])
agent_id = get_first_agent_id(default_dept_token)
print("agent_id", agent_id)
def run_dept_flow(token: str, dept_label: str):
cfgs = list_configs(token, agent_id)
if not cfgs:
raise RuntimeError(f"{dept_label}: configs should have default created")
default_cfg = next((c for c in cfgs if c.get("is_default")), cfgs[0])
print(dept_label, "default_config", default_cfg["id"], default_cfg["name"])
created = create_config(token, agent_id, f"{dept_label}-测试配置A", set_default=False)
print(dept_label, "created_config", created["id"], created["name"])
dup = create_config(token, agent_id, f"{dept_label}-测试配置A", set_default=False)
print(dept_label, "created_duplicate_config", dup["id"], dup["name"])
if dup["name"] == f"{dept_label}-测试配置A":
raise RuntimeError(f"{dept_label}: duplicate name should be auto-renamed with -副本")
updated_default = set_default(token, agent_id, created["id"])
if not updated_default.get("is_default"):
raise RuntimeError(f"{dept_label}: set_default should mark config as default")
cfgs2 = list_configs(token, agent_id)
defaults = [c for c in cfgs2 if c.get("is_default")]
if len(defaults) != 1:
raise RuntimeError(f"{dept_label}: default must be unique, got {len(defaults)}")
cfg_payload = get_config(token, agent_id, created["id"])
if cfg_payload["id"] != created["id"]:
raise RuntimeError(f"{dept_label}: get config mismatch")
updated = update_config(
token,
agent_id,
created["id"],
{
"system_prompt": f"system_prompt::{dept_label}::updated",
"tools": [],
"knowledges": [],
"mcps": [],
},
)
if (updated.get("config_json") or {}).get("context", {}).get(
"system_prompt"
) != f"system_prompt::{dept_label}::updated":
raise RuntimeError(f"{dept_label}: update did not persist system_prompt")
chat_smoke(token, agent_id, created["id"])
delete_config(token, agent_id, created["id"])
delete_config(token, agent_id, dup["id"])
cfgs3 = list_configs(token, agent_id)
if not cfgs3:
raise RuntimeError(f"{dept_label}: configs should not be empty after delete; default should exist")
run_dept_flow(default_dept_token, "default_dept")
run_dept_flow(test_admin_token, "test_dept")
if super_dept_id is None or super_user_id is None:
raise RuntimeError("superadmin token missing department_id/user_id")
tmp_user_payload = {
"username": f"tmp_user_{int(time.time())}",
"password": "tmp_pass_123",
"role": "user",
"department_id": int(super_dept_id),
}
created_user = _request("POST", "/api/auth/users", token=super_token, json_data=tmp_user_payload)
if created_user.status_code != 200:
raise RuntimeError(f"create tmp user failed: {created_user.status_code} {created_user.text}")
tmp_user = created_user.json()
tmp_user_login = tmp_user["user_id"]
tmp_user_id = tmp_user["id"]
tmp_token, _, _ = login(Account(tmp_user_login, "tmp_pass_123", "tmp_user"))
forbidden = _request("POST", f"/api/chat/agent/{agent_id}/configs", token=tmp_token, json_data={"name": "x"})
assert_forbidden(forbidden, "user create config")
deleted_user = _request("DELETE", f"/api/auth/users/{tmp_user_id}", token=super_token)
if deleted_user.status_code != 200:
raise RuntimeError(f"delete tmp user failed: {deleted_user.status_code} {deleted_user.text}")
print("OK")
if __name__ == "__main__":
try:
main()
except Exception as e:
print("FAILED:", e)
sys.exit(1)

View File

@ -1,20 +0,0 @@
from vllm import LLM, SamplingParams
llm = LLM(model="/hdd/zwj/models/meta-llama/Meta-Llama-3-8B-Instruct")
prompts = [
"Hello, my name is",
"The president of the United States is",
"The capital of France is",
"The future of AI is",
]
sampling_params = SamplingParams(temperature=0.8, top_p=0.95)
outputs = llm.generate(prompts, sampling_params)
# Print the outputs.
for output in outputs:
prompt = output.prompt
generated_text = output.outputs[0].text
print(f"Prompt: {prompt!r}, Generated text: {generated_text!r}")

View File

@ -1,61 +0,0 @@
MODEL_DIR=/data/public/models
PORT=8081
TENSOR_PARALLEL_SIZE=1
export CUDA_VISIBLE_DEVICES="0"
source .venv/bin/activate
if [ -z "$1" ]; then
echo "Error: No argument provided. Please specify a model name."
exit 1
fi
if [ "$1" = "qwen3:32b" ]; then
vllm serve "$MODEL_DIR/Qwen/Qwen3-32B" \
--trust-remote-code \
--device cuda --dtype auto --tensor-parallel-size $TENSOR_PARALLEL_SIZE \
--max_model_len 16384 \
--served-model-name "$1" \
--enable-auto-tool-choice \
--tool-call-parser hermes \
--host 0.0.0.0 --port $PORT
fi
# Qwen/Qwen3-Embedding-0.6B
if [ "$1" = "Qwen3-Embedding-0.6B" ]; then
vllm serve "$MODEL_DIR/Qwen/Qwen3-Embedding-0.6B" --task embed \
--trust-remote-code --max_model_len 4096 \
--device cuda --dtype auto --tensor-parallel-size $TENSOR_PARALLEL_SIZE \
--served-model-name "$1" --host 0.0.0.0 --port $PORT
fi
if [ "$1" = "Qwen3-Reranker-0.6B" ]; then
vllm serve "$MODEL_DIR/Qwen/Qwen3-Reranker-0.6B" --task rerank \
--trust-remote-code \
--device cuda --dtype auto --tensor-parallel-size $TENSOR_PARALLEL_SIZE \
--max_model_len 4096 \
--served-model-name "$1" --host 0.0.0.0 --port $PORT
fi
# https://docs.vllm.ai/en/latest/serving/openai_compatible_server.html#named-arguments
# model 模型路径,以文件夹结尾
# tensor-parallel-size 张量并行副本数即GPU的数量咱这儿只有2张卡
# trust-remote-code 信任远程代码主要是为了防止模型初始化时不能执行仓库中的源码默认值是False
# device 用于执行 vLLM 的设备。可选auto、cuda、neuron、cpu
# gpu-memory-utilization 用于模型推理过程的显存占用比例范围为0到1。例如0.5表示显存利用率为 50%。如果未指定,则将使用默认值 0.9。
# dtype “auto”将对 FP16 和 FP32 型使用 FP16 精度,对 BF16 型使用 BF16 精度。
# “half”指FP16 的“一半”。推荐用于 AWQ 量化模型。
# “float16”与“half”相同。
# “bfloat16”用于在精度和范围之间取得平衡。
# “float”是 FP32 精度的简写。
# “float32”表示 FP32 精度。
# kv-cache-dtype kv 缓存存储的数据类型。如果为“auto”则将使用模型默认的数据类型。CUDA 11.8及以上版本 支持 fp8 =fp8_e4m3 和 fp8_e5m2。ROCm AMD GPU 支持 fp8 =fp8_e4m3
# served-model-name 对外提供的API中的模型名称
# host 监听的网络地址0.0.0.0表示所有网卡的所有IP127.0.0.1表示仅限本机
# port API服务的端口

View File

@ -1,19 +0,0 @@
from openai import OpenAI
# Set OpenAI's API key and API base to use vLLM's API server.
openai_api_key = "EMPTY"
openai_api_base = "http://localhost:8080/v1"
client = OpenAI(
api_key=openai_api_key,
base_url=openai_api_base,
)
chat_response = client.chat.completions.create(
model="llama",
messages=[
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Tell me a joke."},
],
)
print("Chat response:", chat_response)

View File

@ -65,6 +65,7 @@ images=(
"nginx:alpine"
"quay.io/coreos/etcd:v3.5.5"
"postgres:16"
"redis:7-alpine"
)
# Pull each image

View File

@ -1,723 +0,0 @@
import asyncio
import json
import traceback
import uuid
from collections.abc import AsyncIterator
from datetime import UTC, datetime
from typing import Any
from langchain.messages import AIMessage, AIMessageChunk, HumanMessage
from langgraph.types import Command
from yuxi import config as conf
from yuxi.agents import agent_manager
from yuxi.plugins.guard import content_guard
from yuxi.repositories.agent_config_repository import AgentConfigRepository
from yuxi.repositories.conversation_repository import ConversationRepository
from yuxi.storage.postgres.manager import pg_manager
from yuxi.utils.logging_config import logger
from yuxi.utils.question_utils import (
normalize_options as _normalize_interrupt_options,
)
from yuxi.utils.question_utils import (
normalize_questions as _normalize_interrupt_questions,
)
def _build_state_files(attachments: list[dict]) -> dict:
"""将附件列表转换为 StateBackend 格式的 files 字典
StateBackend 期望的格式:
{
"/attachments/file.md": {
"content": ["line1", "line2", ...],
"created_at": "...",
"modified_at": "...",
}
}
"""
files = {}
for attachment in attachments:
if attachment.get("status") != "parsed":
continue
file_path = attachment.get("file_path")
markdown = attachment.get("markdown")
if not file_path or not markdown:
continue
now = datetime.now(UTC).isoformat()
# 将 markdown 内容按行拆分
content_lines = markdown.split("\n")
files[file_path] = {
"content": content_lines,
"created_at": attachment.get("uploaded_at", now),
"modified_at": attachment.get("uploaded_at", now),
}
return files
async def _get_langgraph_messages(agent_instance, config_dict):
graph = await agent_instance.get_graph()
state = await graph.aget_state(config_dict)
if not state or not state.values:
logger.warning("No state found in LangGraph")
return None
return state.values.get("messages", [])
def extract_agent_state(values: dict) -> dict:
"""从 LangGraph state 中提取 agent 状态"""
if not isinstance(values, dict):
return {}
# 直接获取,信任 state 的数据结构
todos = values.get("todos")
result = {
"todos": list(todos)[:20] if todos else [],
"files": values.get("files") or {},
}
return result
async def _get_existing_message_ids(conv_repo: ConversationRepository, thread_id: str) -> set[str]:
existing_messages = await conv_repo.get_messages_by_thread_id(thread_id)
return {
msg.extra_metadata["id"]
for msg in existing_messages
if msg.extra_metadata and "id" in msg.extra_metadata and isinstance(msg.extra_metadata["id"], str)
}
async def _save_ai_message(conv_repo: ConversationRepository, thread_id: str, msg_dict: dict) -> None:
content = msg_dict.get("content", "")
tool_calls_data = msg_dict.get("tool_calls", [])
ai_msg = await conv_repo.add_message_by_thread_id(
thread_id=thread_id,
role="assistant",
content=content,
message_type="text",
extra_metadata=msg_dict,
)
if ai_msg and tool_calls_data:
for tc in tool_calls_data:
await conv_repo.add_tool_call(
message_id=ai_msg.id,
tool_name=tc.get("name", "unknown"),
tool_input=tc.get("args", {}),
status="pending",
langgraph_tool_call_id=tc.get("id"),
)
async def _save_tool_message(conv_repo: ConversationRepository, msg_dict: dict) -> None:
tool_call_id = msg_dict.get("tool_call_id")
content = msg_dict.get("content", "")
if not tool_call_id:
return
if isinstance(content, list):
tool_output = json.dumps(content) if content else ""
else:
tool_output = str(content)
await conv_repo.update_tool_call_output(
langgraph_tool_call_id=tool_call_id,
tool_output=tool_output,
status="success",
)
async def save_partial_message(
conv_repo: ConversationRepository,
thread_id: str,
full_msg=None,
error_message: str | None = None,
error_type: str = "interrupted",
):
try:
extra_metadata = {
"error_type": error_type,
"is_error": True,
"error_message": error_message or f"发生错误: {error_type}",
}
if full_msg:
msg_dict = full_msg.model_dump() if hasattr(full_msg, "model_dump") else {}
content = full_msg.content if hasattr(full_msg, "content") else str(full_msg)
extra_metadata = msg_dict | extra_metadata
else:
content = ""
return await conv_repo.add_message_by_thread_id(
thread_id=thread_id,
role="assistant",
content=content,
message_type="text",
extra_metadata=extra_metadata,
)
except Exception as e:
logger.error(f"Error saving message: {e}")
logger.error(traceback.format_exc())
return None
async def save_messages_from_langgraph_state(
agent_instance,
thread_id: str,
conv_repo: ConversationRepository,
config_dict: dict,
) -> None:
try:
messages = await _get_langgraph_messages(agent_instance, config_dict)
if messages is None:
return
existing_ids = await _get_existing_message_ids(conv_repo, thread_id)
for msg in messages:
msg_dict = msg.model_dump() if hasattr(msg, "model_dump") else {}
msg_type = msg_dict.get("type", "unknown")
if msg_type == "human" or getattr(msg, "id", None) in existing_ids:
continue
if msg_type == "ai":
await _save_ai_message(conv_repo, thread_id, msg_dict)
elif msg_type == "tool":
await _save_tool_message(conv_repo, msg_dict)
except Exception as e:
logger.error(f"Error saving messages from LangGraph state: {e}")
logger.error(traceback.format_exc())
def _extract_interrupt_info(state) -> Any | None:
"""从 LangGraph state 中提取中断信息"""
if hasattr(state, "tasks") and state.tasks:
for task in state.tasks:
if hasattr(task, "interrupts") and task.interrupts:
return task.interrupts[0]
interrupt_data = state.values.get("__interrupt__")
if isinstance(interrupt_data, list) and interrupt_data:
return interrupt_data[0]
return None
def _coerce_interrupt_payload(info: Any) -> dict:
"""将 LangGraph interrupt 对象转换为 dict 结构。"""
if isinstance(info, dict):
return info
payload = getattr(info, "value", None)
if isinstance(payload, dict):
return payload
questions = getattr(info, "questions", None)
question = getattr(info, "question", None)
question_id = getattr(info, "question_id", None)
options = getattr(info, "options", None)
multi_select = getattr(info, "multi_select", None)
allow_other = getattr(info, "allow_other", None)
operation = getattr(info, "operation", None)
source = getattr(info, "source", None)
result: dict[str, Any] = {}
if isinstance(questions, list):
result["questions"] = questions
if isinstance(question, str) and question.strip():
result["question"] = question
if isinstance(question_id, str) and question_id.strip():
result["question_id"] = question_id
if isinstance(options, list):
result["options"] = options
if isinstance(multi_select, bool):
result["multi_select"] = multi_select
if isinstance(allow_other, bool):
result["allow_other"] = allow_other
if isinstance(operation, str) and operation.strip():
result["operation"] = operation
if isinstance(source, str) and source.strip():
result["source"] = source
return result
def _build_ask_user_question_payload(info: Any, thread_id: str) -> dict[str, Any]:
"""将 interrupt 信息标准化为 ask_user_question_required 载荷。"""
payload = _coerce_interrupt_payload(info)
questions = _normalize_interrupt_questions(payload.get("questions"))
if not questions:
legacy_question = str(payload.get("question") or "").strip()
if legacy_question:
legacy_item: dict[str, Any] = {
"question_id": str(payload.get("question_id") or uuid.uuid4()),
"question": legacy_question,
"options": _normalize_interrupt_options(payload.get("options")),
"multi_select": bool(payload.get("multi_select", False)),
"allow_other": bool(payload.get("allow_other", True)),
}
legacy_operation = payload.get("operation")
if isinstance(legacy_operation, str) and legacy_operation.strip():
legacy_item["operation"] = legacy_operation.strip()
questions = [legacy_item]
if not questions:
questions = [
{
"question_id": str(uuid.uuid4()),
"question": "请选择一个选项",
"options": [],
"multi_select": False,
"allow_other": True,
}
]
source = str(payload.get("source") or payload.get("tool_name") or "interrupt")
return {
"questions": questions,
"source": source,
"thread_id": thread_id,
}
def _ensure_full_msg(full_msg: AIMessage | None, accumulated_content: list[str]) -> AIMessage | None:
"""如果 full_msg 为空且有累积内容,构建 AIMessage"""
if not full_msg and accumulated_content:
return AIMessage(content="".join(accumulated_content))
return full_msg
async def _resolve_agent_config(
db, agent_id: str, department_id, user_id: str, agent_config_id: int | str | None
) -> tuple:
"""解析 agent_config返回 (config_item, agent_config_id)"""
config_repo = AgentConfigRepository(db)
config_item = None
if agent_config_id is not None:
try:
config_item = await config_repo.get_by_id(int(agent_config_id))
except Exception:
logger.warning(f"Failed to fetch agent config {agent_config_id}: {traceback.format_exc()}")
config_item = None
if config_item is not None and (config_item.department_id != department_id or config_item.agent_id != agent_id):
config_item = None
if config_item is None:
config_item = await config_repo.get_or_create_default(
department_id=department_id, agent_id=agent_id, created_by=user_id
)
agent_config_id = config_item.id
return config_item, agent_config_id
async def check_and_handle_interrupts(
agent,
langgraph_config: dict,
make_chunk,
meta: dict,
thread_id: str,
) -> AsyncIterator[bytes]:
try:
graph = await agent.get_graph()
state = await graph.aget_state(langgraph_config)
if not state or not state.values:
return
interrupt_info = _extract_interrupt_info(state)
if interrupt_info:
question_payload = _build_ask_user_question_payload(interrupt_info, thread_id)
meta["interrupt"] = question_payload
yield make_chunk(status="ask_user_question_required", meta=meta, **question_payload)
except Exception as e:
logger.error(f"Error checking interrupts: {e}")
logger.error(traceback.format_exc())
async def stream_agent_chat(
*,
agent_id: str,
query: str,
config: dict,
meta: dict,
image_content: str | None,
current_user,
db,
) -> AsyncIterator[bytes]:
start_time = asyncio.get_event_loop().time()
def make_chunk(content=None, **kwargs):
return (
json.dumps(
{"request_id": meta.get("request_id"), "response": content, **kwargs}, ensure_ascii=False
).encode("utf-8")
+ b"\n"
)
if image_content:
human_message = HumanMessage(
content=[
{"type": "text", "text": query},
{"type": "image_url", "image_url": {"url": f"data:image/jpeg;base64,{image_content}"}},
]
)
message_type = "multimodal_image"
else:
human_message = HumanMessage(content=query)
message_type = "text"
init_msg = {"role": "user", "content": query, "type": "human"}
if image_content:
init_msg["message_type"] = "multimodal_image"
init_msg["image_content"] = image_content
else:
init_msg["message_type"] = "text"
yield make_chunk(status="init", meta=meta, msg=init_msg)
if conf.enable_content_guard and await content_guard.check(query):
yield make_chunk(
status="error", error_type="content_guard_blocked", error_message="输入内容包含敏感词", meta=meta
)
return
try:
agent = agent_manager.get_agent(agent_id)
except Exception as e:
logger.error(f"Error getting agent {agent_id}: {e}, {traceback.format_exc()}")
yield make_chunk(
status="error",
error_type="agent_error",
error_message=f"智能体 {agent_id} 获取失败: {str(e)}",
meta=meta,
)
return
messages = [human_message]
user_id = str(current_user.id)
department_id = current_user.department_id
if not department_id:
yield make_chunk(status="error", error_type="no_department", error_message="当前用户未绑定部门", meta=meta)
return
agent_config_id = config.get("agent_config_id")
config_item, agent_config_id = await _resolve_agent_config(db, agent_id, department_id, user_id, agent_config_id)
if not (thread_id := config.get("thread_id")):
thread_id = str(uuid.uuid4())
logger.warning(f"No thread_id provided, generated new thread_id: {thread_id}")
agent_config = (config_item.config_json or {}).get("context", {})
input_context = {
"user_id": user_id,
"thread_id": thread_id,
"department_id": department_id,
"agent_config_id": agent_config_id,
"agent_config": agent_config,
}
full_msg = None
accumulated_content: list[str] = []
try:
conv_repo = ConversationRepository(db)
try:
await conv_repo.add_message_by_thread_id(
thread_id=thread_id,
role="user",
content=query,
message_type=message_type,
image_content=image_content,
extra_metadata={"raw_message": human_message.model_dump()},
)
except Exception as e:
logger.error(f"Error saving user message: {e}")
# 先构建 langgraph_config
langgraph_config = {"configurable": {"thread_id": thread_id, "user_id": user_id}}
full_msg = None
accumulated_content = []
async for msg, metadata in agent.stream_messages(messages, input_context=input_context):
if isinstance(msg, AIMessageChunk):
accumulated_content.append(msg.content)
content_for_check = "".join(accumulated_content[-10:])
if conf.enable_content_guard and await content_guard.check_with_keywords(content_for_check):
full_msg = AIMessage(content="".join(accumulated_content))
await save_partial_message(conv_repo, thread_id, full_msg, "content_guard_blocked")
meta["time_cost"] = asyncio.get_event_loop().time() - start_time
yield make_chunk(status="interrupted", message="检测到敏感内容,已中断输出", meta=meta)
return
yield make_chunk(content=msg.content, msg=msg.model_dump(), metadata=metadata, status="loading")
else:
msg_dict = msg.model_dump()
yield make_chunk(msg=msg_dict, metadata=metadata, status="loading")
try:
if msg_dict.get("type") == "tool":
graph = await agent.get_graph()
state = await graph.aget_state(langgraph_config)
agent_state = extract_agent_state(getattr(state, "values", {})) if state else {}
if agent_state:
yield make_chunk(status="agent_state", agent_state=agent_state, meta=meta)
except Exception as e:
logger.error(f"Error processing tool message: {e}")
full_msg = _ensure_full_msg(full_msg, accumulated_content)
if conf.enable_content_guard and hasattr(full_msg, "content") and await content_guard.check(full_msg.content):
await save_partial_message(conv_repo, thread_id, full_msg, "content_guard_blocked")
meta["time_cost"] = asyncio.get_event_loop().time() - start_time
yield make_chunk(status="interrupted", message="检测到敏感内容,已中断输出", meta=meta)
return
async for chunk in check_and_handle_interrupts(agent, langgraph_config, make_chunk, meta, thread_id):
yield chunk
meta["time_cost"] = asyncio.get_event_loop().time() - start_time
try:
graph = await agent.get_graph()
state = await graph.aget_state(langgraph_config)
agent_state = extract_agent_state(getattr(state, "values", {})) if state else {}
except Exception:
agent_state = {}
if agent_state:
yield make_chunk(status="agent_state", agent_state=agent_state, meta=meta)
# 先存储数据库,再返回 finished避免前端查询时数据未落库
await save_messages_from_langgraph_state(
agent_instance=agent,
thread_id=thread_id,
conv_repo=conv_repo,
config_dict=langgraph_config,
)
yield make_chunk(status="finished", meta=meta)
except (asyncio.CancelledError, ConnectionError) as e:
logger.warning(f"Client disconnected, cancelling stream: {e}")
async def save_cleanup():
nonlocal full_msg
full_msg = _ensure_full_msg(full_msg, accumulated_content)
async with pg_manager.get_async_session_context() as new_db:
new_conv_repo = ConversationRepository(new_db)
await save_partial_message(
new_conv_repo,
thread_id,
full_msg=full_msg,
error_message="对话已中断" if not full_msg else None,
error_type="interrupted",
)
cleanup_task = asyncio.create_task(save_cleanup())
try:
await asyncio.shield(cleanup_task)
except asyncio.CancelledError:
pass
except Exception as exc:
logger.error(f"Error during cleanup save: {exc}")
yield make_chunk(status="interrupted", message="对话已中断", meta=meta)
except Exception as e:
logger.error(f"Error streaming messages: {e}, {traceback.format_exc()}")
error_msg = f"Error streaming messages: {e}"
error_type = "unexpected_error"
full_msg = _ensure_full_msg(full_msg, accumulated_content)
async with pg_manager.get_async_session_context() as new_db:
new_conv_repo = ConversationRepository(new_db)
await save_partial_message(
new_conv_repo,
thread_id,
full_msg=full_msg,
error_message=error_msg,
error_type=error_type,
)
yield make_chunk(status="error", error_type=error_type, error_message=error_msg, meta=meta)
async def stream_agent_resume(
*,
agent_id: str,
thread_id: str,
resume_input: Any,
meta: dict,
config: dict,
current_user,
db,
) -> AsyncIterator[bytes]:
start_time = asyncio.get_event_loop().time()
def make_resume_chunk(content=None, **kwargs):
return (
json.dumps(
{"request_id": meta.get("request_id"), "response": content, **kwargs}, ensure_ascii=False
).encode("utf-8")
+ b"\n"
)
try:
agent = agent_manager.get_agent(agent_id)
except Exception as e:
logger.error(f"Error getting agent {agent_id}: {e}, {traceback.format_exc()}")
yield (
f'{{"request_id": "{meta.get("request_id")}", "message": '
f'"Error getting agent {agent_id}: {e}", "status": "error"}}\n'
)
return
init_msg = {"type": "system", "content": f"Resume with input: {resume_input}"}
yield make_resume_chunk(status="init", meta=meta, msg=init_msg)
resume_command = Command(resume=resume_input)
graph = await agent.get_graph()
user_id = str(current_user.id)
department_id = current_user.department_id
if not department_id:
yield make_resume_chunk(
status="error", error_type="no_department", error_message="当前用户未绑定部门", meta=meta
)
return
agent_config_id = (config or {}).get("agent_config_id")
config_item, agent_config_id = await _resolve_agent_config(db, agent_id, department_id, user_id, agent_config_id)
input_context = {
"user_id": user_id,
"thread_id": thread_id,
"department_id": department_id,
"agent_config_id": agent_config_id,
"agent_config": (config_item.config_json or {}).get("context", config_item.config_json or {}),
}
context = agent.context_schema()
agent_config = input_context.get("agent_config")
if isinstance(agent_config, dict):
context.update(agent_config)
context.update(input_context)
stream_source = graph.astream(
resume_command,
context=context,
config={"configurable": {"thread_id": thread_id, "user_id": user_id}},
stream_mode="messages",
)
try:
async for msg, metadata in stream_source:
msg_dict = msg.model_dump()
if "id" not in msg_dict:
msg_dict["id"] = str(uuid.uuid4())
yield make_resume_chunk(
content=getattr(msg, "content", ""), msg=msg_dict, metadata=metadata, status="loading"
)
langgraph_config = {"configurable": {"thread_id": thread_id, "user_id": str(current_user.id)}}
async for chunk in check_and_handle_interrupts(agent, langgraph_config, make_resume_chunk, meta, thread_id):
yield chunk
meta["time_cost"] = asyncio.get_event_loop().time() - start_time
# 先存储数据库,再返回 finished避免前端查询时数据未落库
conv_repo = ConversationRepository(db)
await save_messages_from_langgraph_state(
agent_instance=agent,
thread_id=thread_id,
conv_repo=conv_repo,
config_dict=langgraph_config,
)
yield make_resume_chunk(status="finished", meta=meta)
except (asyncio.CancelledError, ConnectionError) as e:
logger.warning(f"Client disconnected during resume: {e}")
async with pg_manager.get_async_session_context() as new_db:
new_conv_repo = ConversationRepository(new_db)
await save_partial_message(
new_conv_repo, thread_id, error_message="对话恢复已中断", error_type="resume_interrupted"
)
yield make_resume_chunk(status="interrupted", message="对话恢复已中断", meta=meta)
except Exception as e:
logger.error(f"Error during resume: {e}, {traceback.format_exc()}")
async with pg_manager.get_async_session_context() as new_db:
new_conv_repo = ConversationRepository(new_db)
await save_partial_message(
new_conv_repo, thread_id, error_message=f"Error during resume: {e}", error_type="resume_error"
)
yield make_resume_chunk(message=f"Error during resume: {e}", status="error")
async def get_agent_state_view(
*,
agent_id: str,
thread_id: str,
current_user_id: str,
db,
) -> dict:
if not agent_manager.get_agent(agent_id):
from fastapi import HTTPException
raise HTTPException(status_code=404, detail=f"智能体 {agent_id} 不存在")
conv_repo = ConversationRepository(db)
conversation = await conv_repo.get_conversation_by_thread_id(thread_id)
if not conversation or conversation.user_id != str(current_user_id) or conversation.status == "deleted":
from fastapi import HTTPException
raise HTTPException(status_code=404, detail="对话线程不存在")
agent = agent_manager.get_agent(agent_id)
graph = await agent.get_graph()
langgraph_config = {"configurable": {"user_id": str(current_user_id), "thread_id": thread_id}}
state = await graph.aget_state(langgraph_config)
agent_state = extract_agent_state(getattr(state, "values", {})) if state else {}
# 如果 state 中没有 files从附件构建
# 这确保了上传附件后立即可以在文件列表中看到文件
if not agent_state.get("files") or agent_state["files"] == {}:
try:
attachments = await conv_repo.get_attachments_by_thread_id(thread_id)
logger.info(f"[get_agent_state_view] found {len(attachments)} attachments in DB")
if attachments:
first_status = attachments[0].get("status")
first_has_markdown = bool(attachments[0].get("markdown"))
logger.info(
f"[get_agent_state_view] first attachment status: {first_status}, "
f"has markdown: {first_has_markdown}"
)
files = _build_state_files(attachments)
agent_state["files"] = files
logger.info(f"[get_agent_state_view] Built files from attachments: {len(files)} files")
except Exception as e:
logger.warning(f"Failed to fetch attachments for thread {thread_id}: {e}")
return {"agent_state": agent_state}

View File

@ -1,79 +0,0 @@
"""问题和选项规范化工具"""
import uuid
from typing import Any
def normalize_options(raw_options: Any) -> list[dict[str, str]]:
"""规范化选项列表"""
if not isinstance(raw_options, list):
return []
options: list[dict[str, str]] = []
for item in raw_options:
if isinstance(item, dict):
label = str(item.get("label") or item.get("value") or "").strip()
value = str(item.get("value") or item.get("label") or "").strip()
else:
label = str(item).strip()
value = label
if label and value:
options.append({"label": label, "value": value})
return options
def normalize_questions(raw_questions: Any, default_question_id_prefix: str = "q") -> list[dict[str, Any]]:
"""规范化问题列表"""
if not isinstance(raw_questions, list):
return []
questions: list[dict[str, Any]] = []
for idx, item in enumerate(raw_questions):
if not isinstance(item, dict):
continue
question = str(item.get("question") or "").strip()
if not question:
continue
question_id = str(item.get("question_id") or f"{default_question_id_prefix}-{idx + 1}").strip()
if not question_id:
question_id = str(uuid.uuid4())
normalized_question: dict[str, Any] = {
"question_id": question_id,
"question": question,
"options": normalize_options(item.get("options")),
"multi_select": bool(item.get("multi_select", False)),
"allow_other": bool(item.get("allow_other", True)),
}
operation = item.get("operation")
if isinstance(operation, str) and operation.strip():
normalized_question["operation"] = operation.strip()
questions.append(normalized_question)
return questions
def normalize_legacy_question(raw_question: Any) -> dict[str, Any] | None:
"""规范化单个问题(兼容旧格式)"""
if not raw_question:
return None
question = str(raw_question.get("question") or "").strip()
if not question:
return None
question_id = str(raw_question.get("question_id") or "").strip()
if not question_id:
question_id = str(uuid.uuid4())
return {
"question_id": question_id,
"question": question,
"options": normalize_options(raw_question.get("options")),
"multi_select": bool(raw_question.get("multi_select", False)),
"allow_other": bool(raw_question.get("allow_other", True)),
}

View File

@ -174,7 +174,7 @@ export const agentApi = {
/**
* 恢复被人工审批中断的对话流式响应
* @param {string} agentId - 智能体ID
* @param {Object} data - 恢复数据 { thread_id, answer: { question_id: answer }, approved }
* @param {Object} data - 恢复数据 { thread_id, answer: { question_id: answer }, approved }
* @param {Object} options - 可选参数signal, headers等
* @returns {Promise} - 恢复响应流
*/

View File

@ -1164,8 +1164,8 @@ const startRunStream = async (threadId, runId, afterSeq = '0') => {
}
const approvalStatuses = ['ask_user_question_required', 'human_approval_required']
const isApprovalEvent = approvalStatuses.includes(event) ||
approvalStatuses.includes(payload?.chunk?.status)
const isApprovalEvent =
approvalStatuses.includes(event) || approvalStatuses.includes(payload?.chunk?.status)
if (isApprovalEvent) {
const approvalChunk = payload?.chunk || { status: event, thread_id: threadId }

View File

@ -20,9 +20,7 @@
<div v-if="activeQuestion" class="question-block">
<div class="approval-header">
<h4>
{{ activeQuestionIndex + 1 }}. {{ activeQuestion.question }}
</h4>
<h4>{{ activeQuestionIndex + 1 }}. {{ activeQuestion.question }}</h4>
</div>
<div v-if="activeQuestion.operation" class="approval-operation">
@ -77,7 +75,11 @@
<div class="approval-actions">
<button class="btn btn-reject" @click="handleCancel" :disabled="isProcessing">取消</button>
<button class="btn btn-approve" @click="handlePrimaryAction" :disabled="isPrimaryButtonDisabled">
<button
class="btn btn-approve"
@click="handlePrimaryAction"
:disabled="isPrimaryButtonDisabled"
>
{{ primaryButtonText }}
</button>
</div>
@ -92,7 +94,11 @@
<script setup>
import { computed, ref, watch } from 'vue'
import { isOtherOption, normalizeQuestions, DEFAULT_OTHER_OPTION_VALUE } from '@/utils/questionUtils'
import {
isOtherOption,
normalizeQuestions,
DEFAULT_OTHER_OPTION_VALUE
} from '@/utils/questionUtils'
const props = defineProps({
visible: { type: Boolean, default: false },

View File

@ -28,8 +28,10 @@ const parseApprovedDecision = (answer) => {
if (typeof value === 'boolean') return value
if (typeof value === 'string') {
const normalized = value.trim().toLowerCase()
if (normalized === 'approve' || normalized === 'approved' || normalized === 'true') return true
if (normalized === 'reject' || normalized === 'rejected' || normalized === 'false') return false
if (normalized === 'approve' || normalized === 'approved' || normalized === 'true')
return true
if (normalized === 'reject' || normalized === 'rejected' || normalized === 'false')
return false
}
return null
}

View File

@ -9,8 +9,12 @@ const DEFAULT_OTHER_OPTION_VALUE = '__other__'
*/
export const isOtherOption = (option) => {
if (!option || typeof option !== 'object') return false
const label = String(option.label || '').trim().toLowerCase()
const value = String(option.value || '').trim().toLowerCase()
const label = String(option.label || '')
.trim()
.toLowerCase()
const value = String(option.value || '')
.trim()
.toLowerCase()
return (
value === DEFAULT_OTHER_OPTION_VALUE ||
@ -53,14 +57,16 @@ export const normalizeQuestions = (rawQuestions) => {
const question = String(item.question || '').trim()
if (!question) return null
const questionId = String(item.questionId || item.question_id || '').trim() || `q-${index + 1}`
const questionId =
String(item.questionId || item.question_id || '').trim() || `q-${index + 1}`
const operation = String(item.operation || '').trim()
const allowOther = Boolean(item.allowOther ?? item.allow_other ?? true)
const baseOptions = normalizeOptions(item.options || [])
const hasOtherOption = baseOptions.some((option) => isOtherOption(option))
const options = allowOther && !hasOtherOption
? [...baseOptions, { label: '其他', value: DEFAULT_OTHER_OPTION_VALUE }]
: baseOptions
const options =
allowOther && !hasOtherOption
? [...baseOptions, { label: '其他', value: DEFAULT_OTHER_OPTION_VALUE }]
: baseOptions
return {
questionId,

View File

@ -103,7 +103,11 @@
/>
<!-- 反馈模态框 -->
<FeedbackModalComponent v-if="userStore.isAdmin" ref="feedbackModal" :agent-id="selectedAgentId" />
<FeedbackModalComponent
v-if="userStore.isAdmin"
ref="feedbackModal"
:agent-id="selectedAgentId"
/>
<!-- 自定义更多菜单 -->
<Teleport to="body">
@ -186,7 +190,10 @@ const syncSelectedAgentFromRoute = async () => {
const routeAgentExists = (agents.value || []).some((agent) => agent.id === routeAgentId)
if (!routeAgentExists) {
if (selectedAgentId.value) {
await router.replace({ name: 'AgentCompWithId', params: { agent_id: selectedAgentId.value } })
await router.replace({
name: 'AgentCompWithId',
params: { agent_id: selectedAgentId.value }
})
}
return
}