From 1a914a251712f7d361c0780498c1d2326e9bcdf0 Mon Sep 17 00:00:00 2001 From: zhaoxi Date: Wed, 5 Aug 2026 21:26:13 +0800 Subject: [PATCH] =?UTF-8?q?chore:=20=E8=BF=81=E7=A7=BB=E6=9C=8D=E5=8A=A1?= =?UTF-8?q?=E5=99=A8=E5=89=8D=E5=90=8C=E6=AD=A5=E6=9C=AC=E5=9C=B0=E4=BF=AE?= =?UTF-8?q?=E6=94=B9?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Dockerfile | 4 +- data/plugin/data_analytics/api.py | 20 +- .../data_analytics/toolset/ray_submit.py | 12 +- data/toolset/interactive_toolset/approval.py | 8 +- .../regulatory_toolset/query_task_list.py | 6 +- .../query_workflow_status.py | 8 +- data/toolset/regulatory_toolset/send_file.py | 6 +- docker-compose.yml | 2 - .../adapter/model_adapter/agent_factory.py | 7 +- kilostar/api/__init__.py | 15 +- kilostar/api/agent.py | 84 ++++---- kilostar/api/auth.py | 22 +-- kilostar/api/chat.py | 67 +++---- kilostar/api/platform/frontend.py | 6 +- kilostar/api/platform/onebot.py | 6 +- kilostar/api/plugin.py | 44 ++--- kilostar/api/provider.py | 14 +- kilostar/api/resource.py | 84 ++++---- kilostar/api/system.py | 14 +- kilostar/api/task.py | 10 +- kilostar/api/workflow.py | 71 ++++--- .../global_state_machine.py | 85 ++++----- .../core/global_state_machine/gsm_snapshot.py | 77 +++----- .../individual_manager.py | 2 +- .../global_state_machine/provider_manager.py | 6 +- .../global_workflow_manager.py | 3 - .../consciousness_node/consciousness_node.py | 30 +-- .../individual/consciousness_node/template.py | 9 +- .../core/individual/control_node/__init__.py | 17 -- .../individual/control_node/control_node.py | 135 ------------- .../core/individual/control_node/template.py | 51 ----- .../core/individual/growth_node/__init__.py | 14 -- .../individual/growth_node/growth_node.py | 14 -- .../regulatory_node/regulatory_node.py | 8 +- .../individual/regulatory_node/template.py | 12 +- kilostar/core/postgres_database/postgres.py | 4 +- .../core/work/workflow/graph_persistence.py | 10 +- .../core/work/workflow/workflow_engine.py | 94 +++++---- kilostar/plugin_runtime/base_organization.py | 20 +- kilostar/plugin_runtime/plugin_manager.py | 41 ++-- kilostar/plugin_runtime/tool_bridge.py | 22 +-- kilostar/utils/access.py | 6 +- kilostar/utils/actor.py | 180 ++++++++++++++++++ kilostar/utils/agent_model.py | 40 ---- kilostar/utils/get_tool.py | 117 ------------ kilostar/utils/mcp_helper.py | 10 +- kilostar/utils/ray_compat.py | 106 ----------- kilostar/utils/ray_hook.py | 142 -------------- kilostar/utils/request_context.py | 30 +-- kilostar/utils/settings.py | 1 - kilostar/worker_cluster/worker_cluster.py | 35 ++-- kilostar/worker_individual/base_individual.py | 17 +- main.py | 139 ++------------ pyproject.toml | 3 +- tests/conftest.py | 60 ++---- tests/unit/test_agent_factory.py | 6 +- tests/unit/test_api_agent_template.py | 8 +- tests/unit/test_api_chat.py | 35 +--- tests/unit/test_api_custom_toolset_auth.py | 30 +-- tests/unit/test_api_health.py | 10 +- tests/unit/test_api_onebot.py | 10 +- tests/unit/test_api_workflow_auth.py | 12 +- tests/unit/test_gsm_registries.py | 68 ++++--- tests/unit/test_gsm_snapshot.py | 153 ++++----------- tests/unit/test_individual_nodes.py | 4 +- tests/unit/test_plugin_runtime.py | 8 +- tests/unit/test_provider_manager.py | 15 +- tests/unit/test_ray_compat.py | 106 ----------- tests/unit/test_request_context.py | 17 +- tests/unit/test_utils_actor.py | 94 +++++++++ tests/unit/test_utils_get_tool.py | 44 ----- tests/unit/test_utils_ray_hook.py | 107 ----------- tests/unit/test_workflow_engine.py | 92 +++------ 73 files changed, 908 insertions(+), 1961 deletions(-) delete mode 100644 kilostar/core/individual/control_node/__init__.py delete mode 100644 kilostar/core/individual/control_node/control_node.py delete mode 100644 kilostar/core/individual/control_node/template.py delete mode 100644 kilostar/core/individual/growth_node/__init__.py delete mode 100644 kilostar/core/individual/growth_node/growth_node.py create mode 100644 kilostar/utils/actor.py delete mode 100644 kilostar/utils/agent_model.py delete mode 100644 kilostar/utils/get_tool.py delete mode 100644 kilostar/utils/ray_compat.py delete mode 100644 kilostar/utils/ray_hook.py delete mode 100644 tests/unit/test_ray_compat.py create mode 100644 tests/unit/test_utils_actor.py delete mode 100644 tests/unit/test_utils_get_tool.py delete mode 100644 tests/unit/test_utils_ray_hook.py diff --git a/Dockerfile b/Dockerfile index 1928edf..42421f5 100644 --- a/Dockerfile +++ b/Dockerfile @@ -51,8 +51,8 @@ COPY --from=frontend-builder /app/frontend/dist /app/frontend/dist # 重型插件前端 build 产物(让 /plugin-ui// 静态挂载有内容可挂) COPY --from=frontend-builder /app/data/plugin /app/data/plugin -# Expose FastAPI and Ray Dashboard ports -EXPOSE 8000 8265 +# Expose FastAPI port +EXPOSE 8000 # Start the application CMD ["uv", "run", "python", "main.py"] diff --git a/data/plugin/data_analytics/api.py b/data/plugin/data_analytics/api.py index 473dc37..82a38aa 100644 --- a/data/plugin/data_analytics/api.py +++ b/data/plugin/data_analytics/api.py @@ -14,7 +14,7 @@ from fastapi.responses import StreamingResponse from pydantic import BaseModel, Field from kilostar.utils.access import Accessor, TokenData -from kilostar.utils.ray_hook import ray_actor_hook +from kilostar.utils.actor import get_actor router = APIRouter(tags=["data_analytics"]) @@ -40,7 +40,7 @@ class JobCreate(BaseModel): def _get_org(): try: - return ray_actor_hook("org_data_analytics").org_data_analytics + return get_actor("org_data_analytics") except Exception as e: raise HTTPException(503, f"data_analytics 插件未就绪:{e}") @@ -53,7 +53,7 @@ async def list_credentials( token_data: TokenData = Depends(Accessor.get_current_user), ): org = _get_org() - rows = await org.cred_list.remote(token_data.username) + rows = await org.cred_list(token_data.username) return {"credentials": rows} @@ -63,7 +63,7 @@ async def create_credential( token_data: TokenData = Depends(Accessor.get_current_user), ): org = _get_org() - row = await org.cred_create.remote( + row = await org.cred_create( user_id=token_data.username, display_name=payload.display_name, access_key=payload.access_key, @@ -80,7 +80,7 @@ async def delete_credential( token_data: TokenData = Depends(Accessor.get_current_user), ): org = _get_org() - ok = await org.cred_delete.remote(cred_id, token_data.username) + ok = await org.cred_delete(cred_id, token_data.username) if not ok: raise HTTPException(404, "凭证不存在或不属于当前用户") return {"status": "ok"} @@ -96,7 +96,7 @@ async def create_job( ): org = _get_org() try: - return await org.job_create.remote( + return await org.job_create( user_id=token_data.username, cred_id=payload.cred_id, description=payload.description, @@ -110,7 +110,7 @@ async def list_jobs( token_data: TokenData = Depends(Accessor.get_current_user), ): org = _get_org() - rows = await org.job_list.remote(token_data.username) + rows = await org.job_list(token_data.username) return {"jobs": rows} @@ -120,7 +120,7 @@ async def get_job( token_data: TokenData = Depends(Accessor.get_current_user), ): org = _get_org() - row = await org.job_get.remote(job_id, token_data.username) + row = await org.job_get(job_id, token_data.username) if row is None: raise HTTPException(404, "任务不存在") return row @@ -135,7 +135,7 @@ async def stream_job( import json org = _get_org() - row = await org.job_get.remote(job_id, token_data.username) + row = await org.job_get(job_id, token_data.username) if row is None: raise HTTPException(404, "任务不存在") org_task_id = row.get("org_task_id") @@ -143,7 +143,7 @@ async def stream_job( raise HTTPException(409, "任务尚未投递到 organization") async def _generate(): - async for event in await org.stream.remote(org_task_id): + async for event in await org.stream(org_task_id): payload = event if isinstance(event, str) else json.dumps(event, ensure_ascii=False) yield f"data: {payload}\n\n" diff --git a/data/plugin/data_analytics/toolset/ray_submit.py b/data/plugin/data_analytics/toolset/ray_submit.py index 39099f8..b426552 100644 --- a/data/plugin/data_analytics/toolset/ray_submit.py +++ b/data/plugin/data_analytics/toolset/ray_submit.py @@ -1,7 +1,8 @@ -"""ray_submit:把分析脚本提交到 Ray(distributed)或 subprocess(standalone)执行。 +"""ray_submit:把分析脚本提交到子进程执行。 -凭证以 ``AWS_*`` 环境变量注入子进程,让 boto3/pandas-s3 自然读到。 -脚本走 ``kilostar.utils.sandbox.validate_python_code`` 的静态屏蔽兜底。 +名字沿用历史(原设计走 Ray);当前后端为本地子进程,未来要接远程算力只需替换 +执行器,工具签名不变。凭证以 ``AWS_*`` 环境变量注入子进程,让 boto3/pandas-s3 +自然读到。脚本走 ``kilostar.utils.sandbox.validate_python_code`` 的静态屏蔽兜底。 """ from __future__ import annotations @@ -11,7 +12,6 @@ import os import sys import tempfile -from kilostar.utils.ray_compat import _STANDALONE from kilostar.utils.sandbox import ( CodeViolation, get_python_timeout, @@ -33,7 +33,7 @@ def _build_env(creds) -> dict: async def ray_submit(script: str, timeout: int = 300) -> str: - """提交 Python 脚本到 Ray(分布式)或子进程(单机)执行。 + """提交 Python 脚本到子进程执行。 脚本中可直接 ``import boto3`` 读 S3(凭证已通过环境变量注入);可用 pandas / polars / numpy 等已安装的依赖。**只读**——不要尝试 put/delete。 @@ -83,8 +83,6 @@ async def ray_submit(script: str, timeout: int = 300) -> str: if proc.returncode != 0: result += f"\n[exit code: {proc.returncode}]" result = result.strip() or "(no output)" - if not _STANDALONE: - result = f"[mode: ray-cluster (subprocess)]\n{result}" return result except asyncio.TimeoutError: return f"[Error] ray_submit 执行超时({timeout}s)" diff --git a/data/toolset/interactive_toolset/approval.py b/data/toolset/interactive_toolset/approval.py index 543aa05..4d0cbd7 100644 --- a/data/toolset/interactive_toolset/approval.py +++ b/data/toolset/interactive_toolset/approval.py @@ -1,4 +1,4 @@ -from kilostar.utils.ray_hook import ray_actor_hook +from kilostar.utils.actor import get_actor async def approval(message: str, trace_id: str) -> str: @@ -11,7 +11,7 @@ async def approval(message: str, trace_id: str) -> str: Returns: 用户的审批结果 """ - actor_list = ray_actor_hook("global_workflow_manager") - await actor_list.global_workflow_manager.put_pending.remote(trace_id, message) - reply = await actor_list.global_workflow_manager.get_received.remote(trace_id) + gwm = get_actor("global_workflow_manager") + await gwm.put_pending(trace_id, message) + reply = await gwm.get_received(trace_id) return reply diff --git a/data/toolset/regulatory_toolset/query_task_list.py b/data/toolset/regulatory_toolset/query_task_list.py index e4b9a54..740443f 100644 --- a/data/toolset/regulatory_toolset/query_task_list.py +++ b/data/toolset/regulatory_toolset/query_task_list.py @@ -9,7 +9,7 @@ regulatory_node 用以回答"我之前那份报告呢""昨天那个查询结果 from typing import Any, Dict, List, Optional -from kilostar.utils.ray_hook import ray_actor_hook +from kilostar.utils.actor import get_actor async def query_task_list( @@ -34,8 +34,8 @@ async def query_task_list( "total": int } """ - pg = ray_actor_hook("postgres_database").postgres_database - rows: List[Dict[str, Any]] = await pg.list_tasks_by_user.remote( + pg = get_actor("postgres_database") + rows: List[Dict[str, Any]] = await pg.list_tasks_by_user( user_id=user_id, status=status_filter, limit=limit, diff --git a/data/toolset/regulatory_toolset/query_workflow_status.py b/data/toolset/regulatory_toolset/query_workflow_status.py index b23840e..76bbb3b 100644 --- a/data/toolset/regulatory_toolset/query_workflow_status.py +++ b/data/toolset/regulatory_toolset/query_workflow_status.py @@ -6,7 +6,7 @@ regulatory_node 在与用户对话时,可以借此工具回答"我那个任务 from typing import Any, Dict, List -from kilostar.utils.ray_hook import ray_actor_hook +from kilostar.utils.actor import get_actor async def query_workflow_status(trace_id: str, limit: int = 10) -> Dict[str, Any]: @@ -26,10 +26,10 @@ async def query_workflow_status(trace_id: str, limit: int = 10) -> Dict[str, Any ] } """ - pg = ray_actor_hook("postgres_database").postgres_database + pg = get_actor("postgres_database") - workflow = await pg.get_workflow.remote(trace_id) - events = await pg.query_event_logs.remote(trace_id=trace_id, limit=limit) + workflow = await pg.get_workflow(trace_id) + events = await pg.query_event_logs(trace_id=trace_id, limit=limit) recent: List[Dict[str, Any]] = [] for e in events or []: diff --git a/data/toolset/regulatory_toolset/send_file.py b/data/toolset/regulatory_toolset/send_file.py index 89aaa68..35ef528 100644 --- a/data/toolset/regulatory_toolset/send_file.py +++ b/data/toolset/regulatory_toolset/send_file.py @@ -12,7 +12,7 @@ import re import uuid from pathlib import Path -from kilostar.utils.ray_hook import ray_actor_hook +from kilostar.utils.actor import get_actor from kilostar.utils.settings import get_artifact_dir @@ -56,8 +56,8 @@ async def send_file(filename: str, content: str, trace_id: str = "") -> str: }, ensure_ascii=False, ) - actor_list = ray_actor_hook("global_workflow_manager") - await actor_list.global_workflow_manager.put_pending.remote( + gwm = get_actor("global_workflow_manager") + await gwm.put_pending( trace_id, f"__FILE__{payload}" ) return f"已发送文件: {safe_name}" diff --git a/docker-compose.yml b/docker-compose.yml index 0b11e12..4c604eb 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -19,7 +19,6 @@ services: container_name: kilostar_test ports: - "8000:8000" - - "8265:8265" depends_on: db: condition: service_healthy @@ -31,5 +30,4 @@ services: POSTGRES_DB: kilostar SECRET_KEY: test-secret-key-not-for-production KILOSTAR_SECRET_KEY: test-secret-key-not-for-production - KILOSTAR_MODE: standalone KILOSTAR_ENV: dev diff --git a/kilostar/adapter/model_adapter/agent_factory.py b/kilostar/adapter/model_adapter/agent_factory.py index cdc118e..a114c12 100644 --- a/kilostar/adapter/model_adapter/agent_factory.py +++ b/kilostar/adapter/model_adapter/agent_factory.py @@ -24,8 +24,9 @@ from pydantic_ai.providers.deepseek import DeepSeekProvider from pydantic_ai.providers.google import GoogleProvider from pydantic_ai.toolsets import AbstractToolset +from pydantic import BaseModel + from kilostar.core.global_state_machine.model_provider import Provider -from kilostar.utils.agent_model import ResponseModel, DepsModel from kilostar.utils.error import ModelNotExistError @@ -65,9 +66,9 @@ class AgentFactory: self, provider: Provider, model_id: str, - output_type: ResponseModel, + output_type: type[BaseModel], system_prompt: str, - deps_type: DepsModel, + deps_type: type[BaseModel], agent_name: str, tools: list = None, toolsets: Sequence[AbstractToolset[Any]] = None, diff --git a/kilostar/api/__init__.py b/kilostar/api/__init__.py index ba40cb1..f6d87fd 100644 --- a/kilostar/api/__init__.py +++ b/kilostar/api/__init__.py @@ -13,19 +13,14 @@ # limitations under the License. import os -from typing import Dict -from fastapi import FastAPI, WebSocket, Request +from fastapi import FastAPI, Request from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import FileResponse, JSONResponse from fastapi.staticfiles import StaticFiles -from kilostar.utils.ray_compat import _STANDALONE from kilostar.utils.settings import get_settings -if not _STANDALONE: - from ray import serve - from .agent import agent_router from .auth import auth_router from .system import system_router, system_api_router @@ -221,11 +216,3 @@ else: ) -if not _STANDALONE: - @serve.deployment - @serve.ingress(app) - class KiloStarGateway: - gateway: Dict[str, WebSocket] - - def __init__(self): - self.gateway = {} diff --git a/kilostar/api/agent.py b/kilostar/api/agent.py index 56e6d1d..300f5d9 100644 --- a/kilostar/api/agent.py +++ b/kilostar/api/agent.py @@ -14,7 +14,7 @@ from typing import Union -from kilostar.utils.ray_hook import ray_actor_hook +from kilostar.utils.actor import get_actor from fastapi import APIRouter, Depends, Request from pydantic import BaseModel, field_validator from kilostar.utils.access import Accessor, TokenData, RoleChecker @@ -52,8 +52,8 @@ async def get_system_nodes( _: TokenData = Depends(RoleChecker(allowed_roles=UserAuthority.USER)), ): """返回两大系统节点(regulatory/consciousness)当前的持久化配置。""" - postgres_database = ray_actor_hook("postgres_database").postgres_database - configs = await postgres_database.get_all_system_node_configs.remote() + postgres_database = get_actor("postgres_database") + configs = await postgres_database.get_all_system_node_configs() return {"system_nodes": configs} @@ -64,8 +64,8 @@ async def load_agent( _: TokenData = Depends(RoleChecker(allowed_roles=UserAuthority.SUPER_ADMINISTRATOR)), ): """加载/重载某个系统节点的 Agent:先持久化配置,再调用对应节点 Actor 的 ``create_agent``。""" - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine - postgres_database = ray_actor_hook("postgres_database").postgres_database + global_state_machine = get_actor("global_state_machine") + postgres_database = get_actor("postgres_database") accept_lang = request.headers.get("accept-language", "") if isinstance(agent_register, AgentLocalRegister): @@ -73,7 +73,7 @@ async def load_agent( elif isinstance(agent_register, AgentRegister): try: - await postgres_database.upsert_system_node_config.remote( + await postgres_database.upsert_system_node_config( agent_register.individual_name, agent_register.provider_title, agent_register.model_id, @@ -90,14 +90,14 @@ async def load_agent( # Resolve persona system_prompt from DB persona_prompt = None if agent_register.persona_id: - tpl = await postgres_database.get_template.remote(agent_register.persona_id) + tpl = await postgres_database.get_template(agent_register.persona_id) if tpl: persona_prompt = tpl.system_prompt match scope: case "regulatory_node": - node = ray_actor_hook("regulatory_node").regulatory_node - await node.create_agent.remote( + node = get_actor("regulatory_node") + await node.create_agent( global_state_machine, agent_register.provider_title, agent_register.model_id, @@ -107,8 +107,8 @@ async def load_agent( persona_prompt, ) case "consciousness_node": - node = ray_actor_hook("consciousness_node").consciousness_node - await node.create_agent.remote( + node = get_actor("consciousness_node") + await node.create_agent( global_state_machine, agent_register.provider_title, agent_register.model_id, @@ -181,10 +181,10 @@ async def create_worker_individual( token_data: TokenData = Depends(RoleChecker(allowed_roles=UserAuthority.USER)), ): """创建一个 Worker Agent,``owner_id`` 自动绑定为当前登录用户。""" - postgres_database = ray_actor_hook("postgres_database").postgres_database + postgres_database = get_actor("postgres_database") data_dict = worker_data.model_dump() data_dict["owner_id"] = token_data.user_id - worker = await postgres_database.add_worker_individual.remote(**data_dict) + worker = await postgres_database.add_worker_individual(**data_dict) return {"message": "success", "agent_id": worker.agent_id} @@ -197,11 +197,11 @@ async def get_worker_individual_list( plugin_owned slot 是插件登记的"占位 agent",所有用户共享同一份配置, 在前端展示时会用徽标标记,并允许任何登录用户装配 provider/model。 """ - postgres_database = ray_actor_hook("postgres_database").postgres_database - workers = await postgres_database.get_worker_individual_list.remote( + postgres_database = get_actor("postgres_database") + workers = await postgres_database.get_worker_individual_list( owner_id=token_data.user_id ) or [] - all_workers = await postgres_database.get_all_worker_individual.remote() or [] + all_workers = await postgres_database.get_all_worker_individual() or [] seen_ids = {w.agent_id for w in workers} plugin_slots = [ w for w in all_workers @@ -215,8 +215,8 @@ async def get_worker_individual( agent_id: str, token_data: TokenData = Depends(Accessor.get_current_user) ): """按 ``agent_id`` 查询 Worker Agent;非本人的 Agent 返回 403(plugin_owned slot 例外)。""" - postgres_database = ray_actor_hook("postgres_database").postgres_database - worker = await postgres_database.get_worker_individual.remote(agent_id=agent_id) + postgres_database = get_actor("postgres_database") + worker = await postgres_database.get_worker_individual(agent_id=agent_id) if not worker: raise HTTPException(status_code=404, detail="Agent not found") if not getattr(worker, "plugin_owned", None) and worker.owner_id != token_data.user_id: @@ -233,8 +233,8 @@ async def update_worker_individual( token_data: TokenData = Depends(Accessor.get_current_user), ): """局部更新 Worker Agent 配置;同时把状态机里的旧实例移除等待懒加载。""" - postgres_database = ray_actor_hook("postgres_database").postgres_database - worker = await postgres_database.get_worker_individual.remote(agent_id=agent_id) + postgres_database = get_actor("postgres_database") + worker = await postgres_database.get_worker_individual(agent_id=agent_id) if not worker: raise HTTPException(status_code=404, detail="Agent not found") # plugin_owned slot:任何登录用户都能装配 provider/model;普通 worker 仅 owner 可改 @@ -244,21 +244,21 @@ async def update_worker_individual( ) update_data = worker_data.model_dump(exclude_unset=True) - updated_worker = await postgres_database.update_worker_individual.remote( + updated_worker = await postgres_database.update_worker_individual( agent_id=agent_id, **update_data ) - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine + global_state_machine = get_actor("global_state_machine") try: - await global_state_machine.remove_individual.remote(agent_id) + await global_state_machine.remove_individual(agent_id) except Exception: pass # plugin_owned 时顺带触发对应插件的 reload,让新 provider/model 立刻生效 if getattr(worker, "plugin_owned", None): try: - pm = ray_actor_hook("global_plugin_manager").global_plugin_manager - await pm.reload.remote(worker.plugin_owned) + pm = get_actor("global_plugin_manager") + await pm.reload(worker.plugin_owned) except Exception: pass @@ -270,8 +270,8 @@ async def reload_worker_individual( agent_id: str, token_data: TokenData = Depends(Accessor.get_current_user) ): """强制把 Worker 从内存池中卸载,下次调用时按最新配置重新加载。""" - postgres_database = ray_actor_hook("postgres_database").postgres_database - worker = await postgres_database.get_worker_individual.remote(agent_id=agent_id) + postgres_database = get_actor("postgres_database") + worker = await postgres_database.get_worker_individual(agent_id=agent_id) if not worker: raise HTTPException(status_code=404, detail="Agent not found") if worker.owner_id != token_data.user_id: @@ -279,8 +279,8 @@ async def reload_worker_individual( status_code=403, detail="Forbidden: You do not own this agent" ) - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine - await global_state_machine.remove_individual.remote(agent_id) + global_state_machine = get_actor("global_state_machine") + await global_state_machine.remove_individual(agent_id) return {"message": "Worker will be reloaded on next use"} @@ -290,15 +290,15 @@ async def delete_worker_individual( agent_id: str, token_data: TokenData = Depends(Accessor.get_current_user) ): """删除 Worker Agent;非本人 Agent 返回 403。""" - postgres_database = ray_actor_hook("postgres_database").postgres_database - worker = await postgres_database.get_worker_individual.remote(agent_id=agent_id) + postgres_database = get_actor("postgres_database") + worker = await postgres_database.get_worker_individual(agent_id=agent_id) if not worker: raise HTTPException(status_code=404, detail="Agent not found") if worker.owner_id != token_data.user_id: raise HTTPException( status_code=403, detail="Forbidden: You do not own this agent" ) - await postgres_database.delete_worker_individual.remote(agent_id=agent_id) + await postgres_database.delete_worker_individual(agent_id=agent_id) return {"message": "success"} @@ -318,8 +318,8 @@ class PersonaTemplateUpdate(BaseModel): async def list_templates( token_data: TokenData = Depends(Accessor.get_current_user), ): - postgres_database = ray_actor_hook("postgres_database").postgres_database - templates = await postgres_database.list_templates.remote( + postgres_database = get_actor("postgres_database") + templates = await postgres_database.list_templates( owner_id=token_data.user_id ) return {"templates": templates} @@ -330,8 +330,8 @@ async def create_template( data: PersonaTemplateCreate, token_data: TokenData = Depends(RoleChecker(allowed_roles=UserAuthority.USER)), ): - postgres_database = ray_actor_hook("postgres_database").postgres_database - tpl = await postgres_database.add_template.remote( + postgres_database = get_actor("postgres_database") + tpl = await postgres_database.add_template( name=data.name, system_prompt=data.system_prompt, owner_id=token_data.user_id, @@ -345,13 +345,13 @@ async def update_template( data: PersonaTemplateUpdate, token_data: TokenData = Depends(Accessor.get_current_user), ): - postgres_database = ray_actor_hook("postgres_database").postgres_database - tpl = await postgres_database.get_template.remote(template_id) + postgres_database = get_actor("postgres_database") + tpl = await postgres_database.get_template(template_id) if not tpl: raise HTTPException(status_code=404, detail="Template not found") if tpl.owner_id != token_data.user_id: raise HTTPException(status_code=403, detail="Forbidden") - updated = await postgres_database.update_template.remote( + updated = await postgres_database.update_template( template_id, **data.model_dump(exclude_unset=True) ) return {"message": "success", "template": updated} @@ -362,11 +362,11 @@ async def delete_template( template_id: str, token_data: TokenData = Depends(Accessor.get_current_user), ): - postgres_database = ray_actor_hook("postgres_database").postgres_database - tpl = await postgres_database.get_template.remote(template_id) + postgres_database = get_actor("postgres_database") + tpl = await postgres_database.get_template(template_id) if not tpl: raise HTTPException(status_code=404, detail="Template not found") if tpl.owner_id != token_data.user_id: raise HTTPException(status_code=403, detail="Forbidden") - await postgres_database.delete_template.remote(template_id) + await postgres_database.delete_template(template_id) return {"message": "success"} diff --git a/kilostar/api/auth.py b/kilostar/api/auth.py index 8312cc1..a9ad63d 100644 --- a/kilostar/api/auth.py +++ b/kilostar/api/auth.py @@ -17,7 +17,7 @@ from fastapi import Depends from pydantic import BaseModel from kilostar.utils.access import Accessor, TokenData, RoleChecker from fastapi.concurrency import run_in_threadpool -from kilostar.utils.ray_hook import ray_actor_hook +from kilostar.utils.actor import get_actor from kilostar.core.postgres_database.model import UserAuthority from kilostar.utils.error import UserNotExistError from kilostar.utils.rate_limit import register_limiter, login_limiter @@ -36,7 +36,7 @@ class UserRegister(BaseModel): async def create_user(user_register: UserRegister, request: Request): """注册新用户:异步线程池里做 argon2 哈希,再交由 PostgresDatabase Actor 落库。""" register_limiter.check(request) - postgres_database = ray_actor_hook("postgres_database").postgres_database + postgres_database = get_actor("postgres_database") try: hashed_password = await run_in_threadpool( Accessor.hash_password, user_register.password @@ -47,7 +47,7 @@ async def create_user(user_register: UserRegister, request: Request): status_code=400, content={"code": "password_invalid", "message": str(e)}, ) - user = await postgres_database.add_user.remote( + user = await postgres_database.add_user( user_register.user_name, hashed_password ) return {"message": "success", "user_id": user.user_id} @@ -64,8 +64,8 @@ class UserLogin(BaseModel): async def login_user(user_login: UserLogin, request: Request): """用户登录:查询用户后在线程池中校验口令,校验成功则签发 JWT。""" login_limiter.check(request) - postgres_database = ray_actor_hook("postgres_database").postgres_database - user = await postgres_database.login_user.remote(user_login.user_name) + postgres_database = get_actor("postgres_database") + user = await postgres_database.login_user(user_login.user_name) if not user: raise UserNotExistError() tokens = await run_in_threadpool( @@ -113,8 +113,8 @@ async def change_authority( """ Update a user's authority level. Only accessible by SUPER_ADMINISTRATOR. """ - postgres_database = ray_actor_hook("postgres_database").postgres_database - user = await postgres_database.change_user_authority.remote( + postgres_database = get_actor("postgres_database") + user = await postgres_database.change_user_authority( user_id=request.user_id, new_authority=request.new_authority ) return { @@ -133,8 +133,8 @@ async def get_user_list( """ Get a list of all users. Only accessible by SUPER_ADMINISTRATOR. """ - postgres_database = ray_actor_hook("postgres_database").postgres_database - users = await postgres_database.get_all_users.remote() + postgres_database = get_actor("postgres_database") + users = await postgres_database.get_all_users() return { "users": [ {"user_id": u.user_id, "user_name": u.user_name, "role": u.user_authority} @@ -153,6 +153,6 @@ async def delete_user( """ Delete a user. Only accessible by SUPER_ADMINISTRATOR. """ - postgres_database = ray_actor_hook("postgres_database").postgres_database - await postgres_database.delete_user_by_id.remote(user_id=user_id) + postgres_database = get_actor("postgres_database") + await postgres_database.delete_user_by_id(user_id=user_id) return {"message": "success"} diff --git a/kilostar/api/chat.py b/kilostar/api/chat.py index 1abe3e9..43920e2 100644 --- a/kilostar/api/chat.py +++ b/kilostar/api/chat.py @@ -17,7 +17,7 @@ import asyncio from fastapi import APIRouter, Depends, HTTPException, Request from fastapi.responses import StreamingResponse from pydantic import BaseModel -from kilostar.utils.ray_hook import ray_actor_hook +from kilostar.utils.actor import get_actor from kilostar.utils.access import Accessor, TokenData from kilostar.core.individual.regulatory_node.template import ( MessageRequest, @@ -56,8 +56,8 @@ def _build_message_history(rows) -> list: async def _load_message_history(chat_id: str) -> list: - postgres_database = ray_actor_hook("postgres_database").postgres_database - rows = await postgres_database.list_chat_messages.remote(chat_id=chat_id) + postgres_database = get_actor("postgres_database") + rows = await postgres_database.list_chat_messages(chat_id=chat_id) return _build_message_history(rows or []) @@ -76,14 +76,14 @@ async def _ask_regulatory( *, user_id: str, chat_id: str, message: str, message_history: list | None = None ) -> str | None: """统一封装 chat 入口对 RegulatoryNode 的调用。""" - regulatory_node = ray_actor_hook("regulatory_node").regulatory_node + regulatory_node = get_actor("regulatory_node") payload = MessageRequest( platform="client", user_name=user_id, platform_id=chat_id, message=message, ) - resp: MessageResponse | None = await regulatory_node.working.remote( + resp: MessageResponse | None = await regulatory_node.working( payload, message_history ) return _extract_reply(resp) @@ -103,13 +103,13 @@ async def create_chat_session( request: CreateChatRequest, token_data: TokenData = Depends(Accessor.get_current_user), ): - postgres_database = ray_actor_hook("postgres_database").postgres_database - chat = await postgres_database.create_chat_session.remote( + postgres_database = get_actor("postgres_database") + chat = await postgres_database.create_chat_session( user_id=token_data.user_id, title=request.title ) # 存入用户消息 - await postgres_database.add_chat_message.remote( + await postgres_database.add_chat_message( chat_id=chat.chat_id, message=request.initial_message, message_owner="user" ) @@ -122,7 +122,7 @@ async def create_chat_session( # 存入回复消息 if response_msg: - await postgres_database.add_chat_message.remote( + await postgres_database.add_chat_message( chat_id=chat.chat_id, message=response_msg, message_owner="regulatory_node" ) @@ -133,8 +133,8 @@ async def create_chat_session( async def list_chat_sessions( token_data: TokenData = Depends(Accessor.get_current_user), ): - postgres_database = ray_actor_hook("postgres_database").postgres_database - sessions = await postgres_database.list_chat_sessions.remote( + postgres_database = get_actor("postgres_database") + sessions = await postgres_database.list_chat_sessions( user_id=token_data.user_id ) return {"sessions": sessions} @@ -144,8 +144,8 @@ async def list_chat_sessions( async def get_chat_history( chat_id: str, token_data: TokenData = Depends(Accessor.get_current_user) ): - postgres_database = ray_actor_hook("postgres_database").postgres_database - messages = await postgres_database.list_chat_messages.remote(chat_id=chat_id) + postgres_database = get_actor("postgres_database") + messages = await postgres_database.list_chat_messages(chat_id=chat_id) return {"messages": messages} @@ -155,10 +155,10 @@ async def send_chat_message( request: SendMessageRequest, token_data: TokenData = Depends(Accessor.get_current_user), ): - postgres_database = ray_actor_hook("postgres_database").postgres_database + postgres_database = get_actor("postgres_database") # 先取历史(不含当前输入),再写入用户消息,避免历史里出现重复 message_history = await _load_message_history(chat_id) - await postgres_database.add_chat_message.remote( + await postgres_database.add_chat_message( chat_id=chat_id, message=request.message, message_owner="user" ) @@ -172,7 +172,7 @@ async def send_chat_message( # 存回复 if response_msg: - await postgres_database.add_chat_message.remote( + await postgres_database.add_chat_message( chat_id=chat_id, message=response_msg, message_owner="regulatory_node" ) @@ -184,13 +184,13 @@ async def delete_chat_session( chat_id: str, token_data: TokenData = Depends(Accessor.get_current_user), ): - postgres_database = ray_actor_hook("postgres_database").postgres_database - session = await postgres_database.get_chat_session.remote(chat_id=chat_id) + postgres_database = get_actor("postgres_database") + session = await postgres_database.get_chat_session(chat_id=chat_id) if not session: raise HTTPException(status_code=404, detail="Chat session not found") if session.user_id != token_data.user_id: raise HTTPException(status_code=403, detail="Forbidden") - await postgres_database.delete_chat_session.remote(chat_id=chat_id) + await postgres_database.delete_chat_session(chat_id=chat_id) return {"message": "success"} @@ -201,14 +201,12 @@ async def stream_chat_message( request: Request, token_data: TokenData = Depends(Accessor.get_current_user), ): - """SSE 流式聊天端点:standalone 模式下逐 token 流式输出;distributed 模式 fallback 到整段回复。""" - from kilostar.utils.ray_compat import _STANDALONE - - postgres_database = ray_actor_hook("postgres_database").postgres_database + """SSE 流式聊天端点:逐 token 流式输出。""" + postgres_database = get_actor("postgres_database") message_history = await _load_message_history(chat_id) - await postgres_database.add_chat_message.remote( + await postgres_database.add_chat_message( chat_id=chat_id, message=request_body.message, message_owner="user" ) @@ -219,23 +217,12 @@ async def stream_chat_message( message=request_body.message, ) - regulatory_node = ray_actor_hook("regulatory_node").regulatory_node - - if not _STANDALONE: - async def fallback_generator(): - resp = await regulatory_node.working.remote(payload, message_history) - full_response = resp.reply_message if resp else "" - if full_response: - await postgres_database.add_chat_message.remote( - chat_id=chat_id, message=full_response, message_owner="regulatory_node" - ) - yield f"data: {json.dumps({'token': full_response})}\n\n" - yield f"data: {json.dumps({'done': True, 'full_message': full_response})}\n\n" - - return StreamingResponse(fallback_generator(), media_type="text/event-stream") + regulatory_node = get_actor("regulatory_node") token_queue = asyncio.Queue() - stream_task = regulatory_node.stream_working.remote(payload, token_queue, message_history) + stream_task = asyncio.create_task( + regulatory_node.stream_working(payload, token_queue, message_history) + ) async def event_generator(): full_response = "" @@ -260,7 +247,7 @@ async def stream_chat_message( yield f"data: {json.dumps({'token': full_response})}\n\n" if full_response: - await postgres_database.add_chat_message.remote( + await postgres_database.add_chat_message( chat_id=chat_id, message=full_response, message_owner="regulatory_node", diff --git a/kilostar/api/platform/frontend.py b/kilostar/api/platform/frontend.py index 2697ad4..0b5a85f 100644 --- a/kilostar/api/platform/frontend.py +++ b/kilostar/api/platform/frontend.py @@ -15,7 +15,7 @@ from fastapi import APIRouter, Depends, HTTPException, UploadFile, File from pydantic import BaseModel from kilostar.utils.access import Accessor, TokenData -from kilostar.utils.ray_hook import ray_actor_hook +from kilostar.utils.actor import get_actor from kilostar.core.individual.regulatory_node.template import ( MessageRequest, MessageResponse, @@ -42,14 +42,14 @@ async def create_message( """把前端消息转交给 RegulatoryNode 处理,并把回复透传给前端。""" logger.info("收到消息,来源:客户端") logger.debug(f"消息内容:{message.message}") - regulatory_node = ray_actor_hook("regulatory_node").regulatory_node + regulatory_node = get_actor("regulatory_node") msg_request = MessageRequest( platform="client", user_name=token_data.username, platform_id=token_data.user_id, message=message.message, ) - result = await regulatory_node.working.remote(msg_request) + result = await regulatory_node.working(msg_request) if isinstance(result, MessageResponse): return {"message": result.reply_message} if isinstance(result, str): diff --git a/kilostar/api/platform/onebot.py b/kilostar/api/platform/onebot.py index da7f065..e5fabdb 100644 --- a/kilostar/api/platform/onebot.py +++ b/kilostar/api/platform/onebot.py @@ -39,7 +39,7 @@ from kilostar.core.individual.regulatory_node.template import ( MessageResponse, ) from kilostar.utils.logger import get_logger -from kilostar.utils.ray_hook import ray_actor_hook +from kilostar.utils.actor import get_actor logger = get_logger("onebot") @@ -129,8 +129,8 @@ async def _dispatch_event(payload: Dict[str, Any]) -> Optional[Dict[str, Any]]: ) try: - regulatory_node = ray_actor_hook("regulatory_node").regulatory_node - result = await regulatory_node.working.remote(msg_request) + regulatory_node = get_actor("regulatory_node") + result = await regulatory_node.working(msg_request) except Exception as e: logger.exception(f"[OneBot] RegulatoryNode 调用失败: {e}") return None diff --git a/kilostar/api/plugin.py b/kilostar/api/plugin.py index abef5a8..f47fa59 100644 --- a/kilostar/api/plugin.py +++ b/kilostar/api/plugin.py @@ -8,7 +8,7 @@ from fastapi.responses import StreamingResponse from pydantic import BaseModel from kilostar.utils.access import Accessor, TokenData -from kilostar.utils.ray_hook import ray_actor_hook +from kilostar.utils.actor import get_actor from kilostar.utils.settings import get_plugin_dir plugin_router = APIRouter(prefix="/api/v1/plugin", tags=["plugin"]) @@ -25,15 +25,16 @@ async def submit_task( req: SubmitRequest, token_data: TokenData = Depends(Accessor.get_current_user), ): - pm = ray_actor_hook("global_plugin_manager").global_plugin_manager - plugins = await pm.list_plugins.remote() - if req.org_name not in plugins: + pm = get_actor("global_plugin_manager") + plugins = await pm.list_plugins() + # list_plugins 返回 [{"name": ..., ...}],按 name 判断插件是否已加载 + if req.org_name not in {p["name"] for p in plugins}: raise HTTPException(404, f"Plugin '{req.org_name}' not found") - org = ray_actor_hook(f"org_{req.org_name}").get(f"org_{req.org_name}") + org = get_actor(f"org_{req.org_name}") ctx = req.context or {} ctx["user"] = token_data.username - task_id = await org.submit.remote(req.task_description, ctx) + task_id = await org.submit(req.task_description, ctx) return {"task_id": task_id} @@ -42,8 +43,8 @@ async def get_task_status( task_id: str, token_data: TokenData = Depends(Accessor.get_current_user), ): - db = ray_actor_hook("postgres_database").postgres_database - task = await db.get_org_task.remote(task_id) + db = get_actor("postgres_database") + task = await db.get_org_task(task_id) if not task: raise HTTPException(404, "Task not found") return task @@ -54,8 +55,8 @@ async def get_task_events( task_id: str, token_data: TokenData = Depends(Accessor.get_current_user), ): - db = ray_actor_hook("postgres_database").postgres_database - events = await db.query_org_events.remote(task_id) + db = get_actor("postgres_database") + events = await db.query_org_events(task_id) return {"events": events} @@ -64,19 +65,16 @@ async def stream_task( task_id: str, token_data: TokenData = Depends(Accessor.get_current_user), ): - import asyncio - - org_name = None - db = ray_actor_hook("postgres_database").postgres_database - task = await db.get_org_task.remote(task_id) + db = get_actor("postgres_database") + task = await db.get_org_task(task_id) if not task: raise HTTPException(404, "Task not found") org_name = task["org_name"] - org = ray_actor_hook(f"org_{org_name}").get(f"org_{org_name}") + org = get_actor(f"org_{org_name}") async def _generate(): - async for event in await org.stream.remote(task_id): + async for event in await org.stream(task_id): yield f"data: {event}\n\n" return StreamingResponse(_generate(), media_type="text/event-stream") @@ -86,8 +84,8 @@ async def stream_task( async def list_plugins( token_data: TokenData = Depends(Accessor.get_current_user), ): - pm = ray_actor_hook("global_plugin_manager").global_plugin_manager - plugins = await pm.list_plugins.remote() + pm = get_actor("global_plugin_manager") + plugins = await pm.list_plugins() return {"plugins": plugins} @@ -96,8 +94,8 @@ async def install_plugin( name: str, token_data: TokenData = Depends(Accessor.get_current_user), ): - pm = ray_actor_hook("global_plugin_manager").global_plugin_manager - await pm.install.remote(name) + pm = get_actor("global_plugin_manager") + await pm.install(name) return {"status": "ok", "name": name} @@ -106,8 +104,8 @@ async def reload_plugin( name: str, token_data: TokenData = Depends(Accessor.get_current_user), ): - pm = ray_actor_hook("global_plugin_manager").global_plugin_manager - await pm.reload.remote(name) + pm = get_actor("global_plugin_manager") + await pm.reload(name) return {"status": "ok", "name": name} diff --git a/kilostar/api/provider.py b/kilostar/api/provider.py index 9c5737d..e3a4e70 100644 --- a/kilostar/api/provider.py +++ b/kilostar/api/provider.py @@ -18,7 +18,7 @@ from typing import Any, Dict, Literal, Optional from kilostar.utils.access import TokenData, Accessor, RoleChecker from kilostar.core.postgres_database.model import UserAuthority from kilostar.core.global_state_machine.model_provider.base_provider import Provider -from kilostar.utils.ray_hook import ray_actor_hook +from kilostar.utils.actor import get_actor provider_router = APIRouter(prefix="/api/v1/provider", tags=["provider"]) @@ -39,8 +39,8 @@ async def create_provider( token_data: TokenData = Depends(RoleChecker(allowed_roles=UserAuthority.USER)), ) -> None: """注册一个 Provider;owner 为当前登录用户的 ``user_id``。""" - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine - await global_state_machine.add_provider_wrap.remote( + global_state_machine = get_actor("global_state_machine") + await global_state_machine.add_provider_wrap( provider_type=provider_register.provider_type, provider_title=provider_register.provider_title, provider_url=provider_register.provider_url, @@ -61,10 +61,10 @@ async def get_provider_list( _: TokenData = Depends(Accessor.get_current_user), ) -> Dict[str, Any]: """返回当前所有已注册的 Provider,前端用以展示模型清单。apikey 脱敏。""" - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine + global_state_machine = get_actor("global_state_machine") provider_list: Dict[ str, Provider - ] = await global_state_machine.get_provider_list.remote() + ] = await global_state_machine.get_provider_list() masked = {} for title, p in provider_list.items(): d = p.model_dump() if hasattr(p, "model_dump") else dict(p) @@ -138,6 +138,6 @@ async def delete_provider( ), ) -> dict: """删除指定 ``provider_title`` 的 Provider;仅超管可调用。""" - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine - await global_state_machine.delete_provider.remote(provider_title=provider_title) + global_state_machine = get_actor("global_state_machine") + await global_state_machine.delete_provider(provider_title=provider_title) return {"message": "success"} diff --git a/kilostar/api/resource.py b/kilostar/api/resource.py index 8ed8082..29fa9cd 100644 --- a/kilostar/api/resource.py +++ b/kilostar/api/resource.py @@ -15,7 +15,7 @@ from typing import Any, Dict, List, Optional from pydantic import BaseModel import viceroy -from kilostar.utils.ray_hook import ray_actor_hook +from kilostar.utils.actor import get_actor from fastapi import APIRouter, Depends, HTTPException from fastapi.responses import FileResponse from kilostar.utils.access import TokenData, RoleChecker, Accessor @@ -50,7 +50,7 @@ async def install_skill( skill: Skill, _: TokenData = Depends(RoleChecker(allowed_roles=UserAuthority.USER)) ): """通过 viceroy 把 skill 仓库克隆到 ``data/plugin/skill``,并在状态机中登记。""" - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine + global_state_machine = get_actor("global_state_machine") import os from kilostar.utils.settings import get_plugin_dir @@ -63,7 +63,7 @@ async def install_skill( skill_name = skill.path.split("/")[-1] else: skill_name = skill.repo_url.split("/")[-1] - await global_state_machine.add_skill.remote(skill_name) + await global_state_machine.add_skill(skill_name) return {"message": "创建成功"} @@ -72,8 +72,8 @@ async def get_skills( _: TokenData = Depends(RoleChecker(allowed_roles=UserAuthority.USER)), ): """返回当前状态机中已登记的所有 skill 名称列表。""" - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine - skills = await global_state_machine.get_skill_list.remote() + global_state_machine = get_actor("global_state_machine") + skills = await global_state_machine.get_skill_list() return {"skills": skills} @@ -85,8 +85,8 @@ async def delete_skill( ), ): """从状态机中移除 skill 注册项;不会删除磁盘上的代码文件。""" - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine - await global_state_machine.remove_skill.remote(skill_name) + global_state_machine = get_actor("global_state_machine") + await global_state_machine.remove_skill(skill_name) return {"message": "success"} @@ -98,12 +98,12 @@ async def add_mcp_server( _: TokenData = Depends(RoleChecker(allowed_roles=UserAuthority.SUPER_ADMINISTRATOR)), ): """注册一个 MCP 服务器到全局状态机。""" - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine + global_state_machine = get_actor("global_state_machine") import uuid server_id = str(uuid.uuid4())[:8] cfg_dict = config.model_dump(exclude_none=True) - await global_state_machine.add_mcp_server.remote(server_id, cfg_dict) + await global_state_machine.add_mcp_server(server_id, cfg_dict) return {"server_id": server_id, "message": "MCP server registered"} @@ -112,8 +112,8 @@ async def list_mcp_servers( _: TokenData = Depends(RoleChecker(allowed_roles=UserAuthority.USER)), ): """返回已注册的全部 MCP 服务器配置;env 中的敏感字段脱敏。""" - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine - servers = await global_state_machine.list_mcp_servers.remote() + global_state_machine = get_actor("global_state_machine") + servers = await global_state_machine.list_mcp_servers() for s in servers: if "env" in s and isinstance(s["env"], dict): s["env"] = _mask_config(s["env"]) @@ -126,8 +126,8 @@ async def delete_mcp_server( _: TokenData = Depends(RoleChecker(allowed_roles=UserAuthority.SUPER_ADMINISTRATOR)), ): """从状态机中移除一个 MCP 服务器配置。""" - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine - ok = await global_state_machine.delete_mcp_server.remote(server_id) + global_state_machine = get_actor("global_state_machine") + ok = await global_state_machine.delete_mcp_server(server_id) if not ok: raise HTTPException(status_code=404, detail="MCP server not found") return {"message": "success"} @@ -152,8 +152,8 @@ async def download_artifact( if not artifact_id.isalnum() or len(artifact_id) > 32: raise HTTPException(status_code=400, detail="invalid artifact id") - postgres_database = ray_actor_hook("postgres_database").postgres_database - wf = await postgres_database.get_workflow.remote(trace_id) + postgres_database = get_actor("postgres_database") + wf = await postgres_database.get_workflow(trace_id) if not wf: raise HTTPException(status_code=404, detail="Workflow not found") if getattr(wf, "user_id", None) != token_data.user_id: @@ -187,8 +187,8 @@ async def list_toolset_packages( _: TokenData = Depends(RoleChecker(allowed_roles=UserAuthority.USER)), ): """列出所有磁盘上的工具包(``data/toolset//`` 单元)。""" - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine - packages = await global_state_machine.list_toolset_packages.remote() + global_state_machine = get_actor("global_state_machine") + packages = await global_state_machine.list_toolset_packages() return {"packages": packages} @@ -198,8 +198,8 @@ async def get_toolset_package_readme( _: TokenData = Depends(RoleChecker(allowed_roles=UserAuthority.USER)), ): """返回指定工具包的 README.md 内容(markdown 文本)。""" - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine - content = await global_state_machine.get_toolset_package_readme.remote(name) + global_state_machine = get_actor("global_state_machine") + content = await global_state_machine.get_toolset_package_readme(name) if content is None: raise HTTPException(status_code=404, detail="README not found") return {"name": name, "content": content} @@ -216,9 +216,9 @@ async def get_tools( 其中 ``mcp_servers`` 会现场尝试连接每个已注册的 MCP 服务器并列出它们暴露的 工具名,便于前端展示;任意一台 MCP server 不可达不影响其他工具的返回。 """ - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine - tool_mapper = await global_state_machine.get_tool_mapper.remote() - categories = await global_state_machine.get_tool_categories.remote() + global_state_machine = get_actor("global_state_machine") + tool_mapper = await global_state_machine.get_tool_mapper() + categories = await global_state_machine.get_tool_categories() all_tool_names = set() for scope_tools in tool_mapper.values(): @@ -266,8 +266,8 @@ async def list_tool_configs( _: TokenData = Depends(RoleChecker(allowed_roles=UserAuthority.SUPER_ADMINISTRATOR)), ): """列出所有工具运行期配置;敏感字段会被脱敏。""" - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine - raw = await global_state_machine.list_tool_configs.remote() + global_state_machine = get_actor("global_state_machine") + raw = await global_state_machine.list_tool_configs() return { "configs": {name: _mask_config(cfg) for name, cfg in raw.items()}, } @@ -279,8 +279,8 @@ async def get_tool_config( _: TokenData = Depends(RoleChecker(allowed_roles=UserAuthority.SUPER_ADMINISTRATOR)), ): """按工具名取出脱敏后的配置。""" - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine - raw = await global_state_machine.get_tool_config.remote(tool_name) + global_state_machine = get_actor("global_state_machine") + raw = await global_state_machine.get_tool_config(tool_name) return {"tool_name": tool_name, "config": _mask_config(raw)} @@ -291,8 +291,8 @@ async def set_tool_config( _: TokenData = Depends(RoleChecker(allowed_roles=UserAuthority.SUPER_ADMINISTRATOR)), ): """写入/覆盖某工具的运行期配置(如 ``tavily_search`` 的 ``api_key``)。""" - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine - await global_state_machine.set_tool_config.remote(tool_name, body.config) + global_state_machine = get_actor("global_state_machine") + await global_state_machine.set_tool_config(tool_name, body.config) return {"message": "success"} @@ -302,8 +302,8 @@ async def delete_tool_config( _: TokenData = Depends(RoleChecker(allowed_roles=UserAuthority.SUPER_ADMINISTRATOR)), ): """删除某工具的运行期配置。""" - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine - ok = await global_state_machine.delete_tool_config.remote(tool_name) + global_state_machine = get_actor("global_state_machine") + ok = await global_state_machine.delete_tool_config(tool_name) if not ok: raise HTTPException(status_code=404, detail="Tool config not found") return {"message": "success"} @@ -343,12 +343,12 @@ async def create_custom_toolset( body: CustomToolsetCreate, token_data: TokenData = Depends(RoleChecker(allowed_roles=UserAuthority.USER)), ): - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine + global_state_machine = get_actor("global_state_machine") import uuid toolset_id = str(uuid.uuid4())[:8] try: - saved = await global_state_machine.add_custom_toolset.remote( + saved = await global_state_machine.add_custom_toolset( toolset_id=toolset_id, name=body.name, tools=body.tools, @@ -368,8 +368,8 @@ async def list_custom_toolsets( """列出工具组:支持按 category 过滤。USER 只能看到自己的+系统的;ADMIN 看全部。""" from kilostar.utils.access import get_authority - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine - toolsets = await global_state_machine.list_custom_toolsets.remote() + global_state_machine = get_actor("global_state_machine") + toolsets = await global_state_machine.list_custom_toolsets() authority = await get_authority(token_data.user_id) if authority < UserAuthority.ADMINISTRATOR: toolsets = [ @@ -386,8 +386,8 @@ async def get_custom_toolset( toolset_id: str, token_data: TokenData = Depends(RoleChecker(allowed_roles=UserAuthority.USER)), ): - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine - ts = await global_state_machine.get_custom_toolset.remote(toolset_id) + global_state_machine = get_actor("global_state_machine") + ts = await global_state_machine.get_custom_toolset(toolset_id) if not ts: raise HTTPException(status_code=404, detail="Custom toolset not found") await _assert_toolset_owner_or_admin(ts, token_data) @@ -400,8 +400,8 @@ async def update_custom_toolset( body: CustomToolsetUpdate, token_data: TokenData = Depends(RoleChecker(allowed_roles=UserAuthority.USER)), ): - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine - existing = await global_state_machine.get_custom_toolset.remote(toolset_id) + global_state_machine = get_actor("global_state_machine") + existing = await global_state_machine.get_custom_toolset(toolset_id) if not existing: raise HTTPException(status_code=404, detail="Custom toolset not found") if existing.get("is_system"): @@ -411,7 +411,7 @@ async def update_custom_toolset( tools = body.tools if body.tools is not None else existing["tools"] description = body.description if body.description is not None else existing.get("description") try: - saved = await global_state_machine.add_custom_toolset.remote( + saved = await global_state_machine.add_custom_toolset( toolset_id=toolset_id, name=name, tools=tools, @@ -429,14 +429,14 @@ async def delete_custom_toolset( token_data: TokenData = Depends(RoleChecker(allowed_roles=UserAuthority.USER)), ): """删除工具组:系统预置不可删;USER 只能删自己的;ADMIN 及以上可删任意用户的。""" - global_state_machine = ray_actor_hook("global_state_machine").global_state_machine - existing = await global_state_machine.get_custom_toolset.remote(toolset_id) + global_state_machine = get_actor("global_state_machine") + existing = await global_state_machine.get_custom_toolset(toolset_id) if not existing: raise HTTPException(status_code=404, detail="Custom toolset not found") if existing.get("is_system"): raise HTTPException(status_code=403, detail="系统预置工具集不可删除") await _assert_toolset_owner_or_admin(existing, token_data) - ok = await global_state_machine.delete_custom_toolset.remote(toolset_id) + ok = await global_state_machine.delete_custom_toolset(toolset_id) if not ok: raise HTTPException(status_code=404, detail="Custom toolset not found") return {"message": "success"} diff --git a/kilostar/api/system.py b/kilostar/api/system.py index 5fe0e35..3f795d2 100644 --- a/kilostar/api/system.py +++ b/kilostar/api/system.py @@ -24,7 +24,7 @@ from __future__ import annotations from fastapi import APIRouter, Depends from fastapi.responses import JSONResponse -from kilostar.utils.ray_hook import ray_actor_hook +from kilostar.utils.actor import get_actor from kilostar.utils.access import Accessor, TokenData, RoleChecker from kilostar.core.postgres_database.model import UserAuthority from kilostar.utils.config_loader import ( @@ -49,15 +49,15 @@ async def readiness(): checks = {"postgres": False, "global_state_machine": False} try: - postgres_database = ray_actor_hook("postgres_database").postgres_database - await postgres_database.ping.remote() + postgres_database = get_actor("postgres_database") + await postgres_database.ping() checks["postgres"] = True except Exception: pass try: - gsm = ray_actor_hook("global_state_machine").global_state_machine - await gsm.get_skill_list.remote() + gsm = get_actor("global_state_machine") + await gsm.get_skill_list() checks["global_state_machine"] = True except Exception: pass @@ -95,8 +95,8 @@ async def query_system_logs( offset: int = 0, _: TokenData = Depends(Accessor.get_current_user), ): - pg = ray_actor_hook("postgres_database").postgres_database - logs = await pg.query_event_logs.remote( + pg = get_actor("postgres_database") + logs = await pg.query_event_logs( trace_id=trace_id, event_type=event_type, level=level, diff --git a/kilostar/api/task.py b/kilostar/api/task.py index 5d8cdf1..43890d9 100644 --- a/kilostar/api/task.py +++ b/kilostar/api/task.py @@ -17,7 +17,7 @@ from typing import Optional from fastapi import APIRouter, Depends, HTTPException from kilostar.utils.access import Accessor, TokenData -from kilostar.utils.ray_hook import ray_actor_hook +from kilostar.utils.actor import get_actor task_router = APIRouter(prefix="/api/v1/task", tags=["task"]) @@ -30,8 +30,8 @@ async def list_tasks( token_data: TokenData = Depends(Accessor.get_current_user), ): """列出当前用户的所有短任务,按时间倒序。""" - postgres_database = ray_actor_hook("postgres_database").postgres_database - tasks = await postgres_database.list_tasks_by_user.remote( + postgres_database = get_actor("postgres_database") + tasks = await postgres_database.list_tasks_by_user( user_id=token_data.user_id, status=status, limit=limit, @@ -46,8 +46,8 @@ async def get_task( token_data: TokenData = Depends(Accessor.get_current_user), ): """按 task_id 读取一条 task 详情。仅 owner 可访问。""" - postgres_database = ray_actor_hook("postgres_database").postgres_database - task = await postgres_database.get_task.remote(task_id) + postgres_database = get_actor("postgres_database") + task = await postgres_database.get_task(task_id) if not task: raise HTTPException(status_code=404, detail="task not found") if task.get("user_id") != token_data.user_id: diff --git a/kilostar/api/workflow.py b/kilostar/api/workflow.py index 35a731a..285f1bb 100644 --- a/kilostar/api/workflow.py +++ b/kilostar/api/workflow.py @@ -12,7 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -from kilostar.utils.ray_hook import ray_actor_hook +from kilostar.utils.actor import get_actor, spawn_background from fastapi import APIRouter, Request, HTTPException, Depends from fastapi.responses import StreamingResponse from pydantic import BaseModel @@ -34,22 +34,23 @@ async def create_workflow( request: CreateWorkflowRequest, token_data: TokenData = Depends(Accessor.get_current_user), ): - postgres_database = ray_actor_hook("postgres_database").postgres_database + postgres_database = get_actor("postgres_database") trace_id = str(ULID()) - await postgres_database.create_workflow.remote( + await postgres_database.create_workflow( trace_id=trace_id, user_id=token_data.user_id, title=request.title, command=request.command, ) - global_workflow_manager = ray_actor_hook( - "global_workflow_manager" - ).global_workflow_manager - await global_workflow_manager.create_trace.remote(trace_id) + global_workflow_manager = get_actor("global_workflow_manager") + await global_workflow_manager.create_trace(trace_id) - consciousness_node = ray_actor_hook("consciousness_node").consciousness_node - consciousness_node.start_workflow_design.remote(trace_id, request.command) + consciousness_node = get_actor("consciousness_node") + # fire-and-forget:设计过程可能跑很久,不能阻塞 HTTP 响应 + spawn_background( + consciousness_node.start_workflow_design(trace_id, request.command) + ) return {"trace_id": trace_id, "status": "creating"} @@ -57,8 +58,8 @@ async def create_workflow( @workflow_router.get("/list") async def get_workflow_list( token_data: TokenData = Depends(RoleChecker(allowed_roles=UserAuthority.USER))): - postgres_database = ray_actor_hook("postgres_database").postgres_database - workflows = await postgres_database.list_workflows.remote( + postgres_database = get_actor("postgres_database") + workflows = await postgres_database.list_workflows( user_id=token_data.user_id ) return {"workflows": workflows} @@ -75,23 +76,21 @@ async def get_workflow_sse( 鉴权走标准 ``Authorization: Bearer`` 头(前端用 fetch-based SSE, token 不进 URL)。校验该 trace_id 属于当前用户。 """ - postgres_database = ray_actor_hook("postgres_database").postgres_database - wf = await postgres_database.get_workflow.remote(trace_id) + postgres_database = get_actor("postgres_database") + wf = await postgres_database.get_workflow(trace_id) if not wf: raise HTTPException(status_code=404, detail="Workflow not found") if getattr(wf, "user_id", None) != token_data.user_id: raise HTTPException(status_code=403, detail="Forbidden") - global_workflow_manager = ray_actor_hook( - "global_workflow_manager" - ).global_workflow_manager + global_workflow_manager = get_actor("global_workflow_manager") async def event_generator(): try: while True: if await request.is_disconnected(): break - message = await global_workflow_manager.get_pending.remote(trace_id) + message = await global_workflow_manager.get_pending(trace_id) if message: yield f"data: {message}\n\n" else: @@ -108,8 +107,8 @@ async def post_workflow_reply( request: Request, token_data: TokenData = Depends(Accessor.get_current_user), ): - postgres_database = ray_actor_hook("postgres_database").postgres_database - wf = await postgres_database.get_workflow.remote(trace_id) + postgres_database = get_actor("postgres_database") + wf = await postgres_database.get_workflow(trace_id) if not wf: raise HTTPException(status_code=404, detail="Workflow not found") if getattr(wf, "user_id", None) != token_data.user_id: @@ -117,10 +116,8 @@ async def post_workflow_reply( data = await request.json() reply_msg = data.get("message", "") - global_workflow_manager = ray_actor_hook( - "global_workflow_manager" - ).global_workflow_manager - await global_workflow_manager.put_received.remote(trace_id, reply_msg) + global_workflow_manager = get_actor("global_workflow_manager") + await global_workflow_manager.put_received(trace_id, reply_msg) return {"status": "ok"} @@ -128,14 +125,14 @@ async def post_workflow_reply( async def get_workflow_detail( trace_id: str, token_data: TokenData = Depends(Accessor.get_current_user) ): - postgres_database = ray_actor_hook("postgres_database").postgres_database - wf = await postgres_database.get_workflow.remote(trace_id) + postgres_database = get_actor("postgres_database") + wf = await postgres_database.get_workflow(trace_id) if not wf: raise HTTPException(status_code=404, detail="Workflow not found") if getattr(wf, "user_id", None) != token_data.user_id: raise HTTPException(status_code=403, detail="Forbidden") - context = await postgres_database.get_workflow_context.remote(trace_id) + context = await postgres_database.get_workflow_context(trace_id) work_link = ( context.work_link if context and hasattr(context, "work_link") else [] @@ -209,32 +206,30 @@ async def resume_workflow( ): """从 ``workflow_graph_state`` 持久化恢复一个被中断/挂起的工作流。 - 新 fire 一个 ray task,task 入口的 ``hydrate`` 检查会自动走 resume 路径 + 新 fire 一个后台任务,入口的 ``hydrate`` 检查会自动走 resume 路径 把剩余节点跑完。 """ - postgres_database = ray_actor_hook("postgres_database").postgres_database - wf = await postgres_database.get_workflow.remote(trace_id) + postgres_database = get_actor("postgres_database") + wf = await postgres_database.get_workflow(trace_id) if not wf: raise HTTPException(status_code=404, detail="Workflow not found") if getattr(wf, "user_id", None) != token_data.user_id: raise HTTPException(status_code=403, detail="Forbidden") - record = await postgres_database.get_workflow_graph_state.remote(trace_id) + record = await postgres_database.get_workflow_graph_state(trace_id) if record is None: raise HTTPException( status_code=409, detail="该工作流没有可恢复的图持久化记录" ) - global_workflow_manager = ray_actor_hook( - "global_workflow_manager" - ).global_workflow_manager - await global_workflow_manager.create_trace.remote(trace_id) + global_workflow_manager = get_actor("global_workflow_manager") + await global_workflow_manager.create_trace(trace_id) from kilostar.core.work.workflow.workflow_engine import run_workflow_task # resume_only=True:task 入口 hydrate 失败会 fail-fast,绝不 fall through # 到"全新模式空跑"。workflow_data 在 resume 路径上不会被使用,传空 dict 占位。 - run_workflow_task.remote({}, trace_id, resume_only=True) + spawn_background(run_workflow_task({}, trace_id, resume_only=True)) return {"trace_id": trace_id, "status": "resuming"} @@ -251,15 +246,15 @@ async def get_workflow_graph_mermaid( """ from kilostar.core.work.workflow.workflow_engine import workflow_graph - postgres_database = ray_actor_hook("postgres_database").postgres_database - wf = await postgres_database.get_workflow.remote(trace_id) + postgres_database = get_actor("postgres_database") + wf = await postgres_database.get_workflow(trace_id) if not wf: raise HTTPException(status_code=404, detail="Workflow not found") if getattr(wf, "user_id", None) != token_data.user_id: raise HTTPException(status_code=403, detail="Forbidden") visited: list[str] = [] - record = await postgres_database.get_workflow_graph_state.remote(trace_id) + record = await postgres_database.get_workflow_graph_state(trace_id) if record is not None: history = getattr(record, "history", None) or [] # history 里每条 NodeSnapshot.id 形如 "ClassName:hash",截前缀作为 NodeIdent diff --git a/kilostar/core/global_state_machine/global_state_machine.py b/kilostar/core/global_state_machine/global_state_machine.py index a008c69..a7d5f2e 100644 --- a/kilostar/core/global_state_machine/global_state_machine.py +++ b/kilostar/core/global_state_machine/global_state_machine.py @@ -13,10 +13,6 @@ # limitations under the License. from typing import Any, Dict, List, Optional, Tuple -from kilostar.utils.ray_compat import actor_class, _STANDALONE - -if not _STANDALONE: - import ray from kilostar.core.global_state_machine.individual_manager import ( GlobalIndividualManager, @@ -28,7 +24,6 @@ from kilostar.core.global_state_machine.gsm_snapshot import GSMSnapshot from kilostar.core.postgres_database import PostgresDatabase -@actor_class class GlobalStateMachine: """全局状态机 Actor,统一持有 Provider/Tool/Skill/Individual/MCP/CustomToolset 注册表。 @@ -62,13 +57,13 @@ class GlobalStateMachine: self.postgres_database ) # MCP servers - rows = await self.postgres_database.list_mcp_servers_db.remote() + rows = await self.postgres_database.list_mcp_servers_db() self._mcp_servers = {row["server_id"]: row for row in rows} # Tool configs - cfg_rows = await self.postgres_database.list_tool_configs_db.remote() + cfg_rows = await self.postgres_database.list_tool_configs_db() self._tool_configs = {row["tool_name"]: row["config"] for row in cfg_rows} # Custom toolsets(含系统预置) - ts_rows = await self.postgres_database.list_custom_toolsets.remote() + ts_rows = await self.postgres_database.list_custom_toolsets() self._custom_toolsets = {row["toolset_id"]: row for row in ts_rows} # 补种系统预置工具集(首次启动或 DB 被 create_all 重建后) await self._seed_system_toolsets() @@ -96,12 +91,12 @@ class GlobalStateMachine: continue if tid in wanted_ids: continue - await self.postgres_database.delete_custom_toolset.remote(tid) + await self.postgres_database.delete_custom_toolset(tid) self._custom_toolsets.pop(tid, None) for name, pkg in packages.items(): tid = f"system::{name}" - saved = await self.postgres_database.upsert_custom_toolset.remote( + saved = await self.postgres_database.upsert_custom_toolset( toolset_id=tid, name=pkg.get("display_name") or name, tools=list(pkg.get("tools", [])), @@ -155,18 +150,10 @@ class GlobalStateMachine: def _publish_snapshot(self) -> None: """版本号 +1 并发布当前状态快照。""" self._config_version += 1 - snapshot = self._build_snapshot() - if _STANDALONE: - self._current_ref = snapshot - else: - self._current_ref = ray.put(snapshot) + self._current_ref = self._build_snapshot() async def current_config_ref(self) -> Tuple[int, Any]: - """返回 ``(version, ObjectRef 或 snapshot)``。 - - 分布式模式返回 ObjectRef,调用方用 ``ray.get`` 自取; - 单机模式直接返回 snapshot 对象。 - """ + """返回 ``(version, snapshot)``。进程内直接持有 snapshot 对象。""" if self._current_ref is None: self._publish_snapshot() return self._config_version, self._current_ref @@ -199,11 +186,11 @@ class GlobalStateMachine: self._publish_snapshot() return result - def get_provider_list(self): + async def get_provider_list(self): """返回内存中已登记的全部 Provider。""" return self._global_provider_manager.get_provider_list() - def get_provider(self, provider_title): + async def get_provider(self, provider_title): """按 provider_title 取出单个 Provider 实例。""" return self._global_provider_manager.get_provider(provider_title) @@ -217,17 +204,17 @@ class GlobalStateMachine: # ─── Tool / Toolset ──────────────────────────────────────── - def get_tool_mapper(self): + async def get_tool_mapper(self): """返回 agent_name -> {tool_name: callable} 的全量映射。""" return self._global_tool_manager.tool_mapper - def get_tool_list(self, agent_name: str): + async def get_tool_list(self, agent_name: str): """返回某个 agent 可用的工具集(其专属工具与 default 工具的并集)。""" tools = self._global_tool_manager.tool_mapper.get(agent_name, {}) default_tools = self._global_tool_manager.tool_mapper.get("default", {}) return {**default_tools, **tools} - def get_tool_categories(self): + async def get_tool_categories(self): """返回工具按分类聚合的完整信息。""" return { "system": self._global_tool_manager.get_system_tools(), @@ -241,22 +228,22 @@ class GlobalStateMachine: "all": self._global_tool_manager.get_all_tools(), } - def get_toolsets_for_scope(self, scope: str) -> List[Any]: + async def get_toolsets_for_scope(self, scope: str) -> List[Any]: """返回某个 scope 下的"系统 + 自定义工具组"toolset 列表(不含 MCP)。""" return self._global_tool_manager.get_toolsets_for_scope(scope) - def get_retrieval_toolsets_for_scope(self, scope: str) -> List[Any]: + async def get_retrieval_toolsets_for_scope(self, scope: str) -> List[Any]: """仅返回 retrieval 工具集(system_node 专用,不包含 generation 工具)。""" return self._global_tool_manager.get_retrieval_toolsets_for_scope(scope) - def list_toolset_packages(self) -> List[Dict[str, Any]]: + async def list_toolset_packages(self) -> List[Dict[str, Any]]: """列出所有磁盘工具包(前端"工具插件"页面卡片即由此渲染)。""" return [ {k: v for k, v in pkg.items() if k != "readme_path"} for pkg in self._global_tool_manager.toolset_packages.values() ] - def get_toolset_package_readme(self, name: str) -> Optional[str]: + async def get_toolset_package_readme(self, name: str) -> Optional[str]: """读取指定工具包的 README.md 内容;不存在返回 None。""" pkg = self._global_tool_manager.toolset_packages.get(name) if not pkg or not pkg.get("readme_path"): @@ -271,26 +258,26 @@ class GlobalStateMachine: async def add_mcp_server(self, server_id: str, config: Dict[str, Any]) -> bool: """注册一个 MCP 服务器配置(写库 → 写内存)。""" - saved = await self.postgres_database.upsert_mcp_server.remote(server_id, config) + saved = await self.postgres_database.upsert_mcp_server(server_id, config) self._mcp_servers[server_id] = saved self._publish_snapshot() return True - def get_mcp_server(self, server_id: str) -> Optional[Dict[str, Any]]: + async def get_mcp_server(self, server_id: str) -> Optional[Dict[str, Any]]: return self._mcp_servers.get(server_id) - def list_mcp_servers(self) -> List[Dict[str, Any]]: + async def list_mcp_servers(self) -> List[Dict[str, Any]]: return [ {"server_id": sid, **cfg} for sid, cfg in self._mcp_servers.items() ] async def delete_mcp_server(self, server_id: str) -> bool: - ok = await self.postgres_database.delete_mcp_server_db.remote(server_id) + ok = await self.postgres_database.delete_mcp_server_db(server_id) self._mcp_servers.pop(server_id, None) self._publish_snapshot() return bool(ok) - def get_mcp_server_configs(self) -> Dict[str, Dict[str, Any]]: + async def get_mcp_server_configs(self) -> Dict[str, Dict[str, Any]]: """返回原始 MCP 服务器配置字典(供节点创建 toolsets 时使用)。""" return dict(self._mcp_servers) @@ -298,23 +285,23 @@ class GlobalStateMachine: async def set_tool_config(self, tool_name: str, config: Dict[str, Any]) -> bool: """整体覆盖某工具的运行期配置(敏感字段在 DAO 内自动加密)。""" - saved = await self.postgres_database.upsert_tool_config.remote( + saved = await self.postgres_database.upsert_tool_config( tool_name, config ) self._tool_configs[tool_name] = saved["config"] self._publish_snapshot() return True - def get_tool_config(self, tool_name: str) -> Dict[str, Any]: + async def get_tool_config(self, tool_name: str) -> Dict[str, Any]: """按工具名取出配置;不存在则返回空字典。""" return dict(self._tool_configs.get(tool_name, {})) - def list_tool_configs(self) -> Dict[str, Dict[str, Any]]: + async def list_tool_configs(self) -> Dict[str, Dict[str, Any]]: """返回全部已配置工具的配置(包含敏感字段,调用方需自行脱敏)。""" return dict(self._tool_configs) async def delete_tool_config(self, tool_name: str) -> bool: - ok = await self.postgres_database.delete_tool_config_db.remote(tool_name) + ok = await self.postgres_database.delete_tool_config_db(tool_name) self._tool_configs.pop(tool_name, None) self._publish_snapshot() return bool(ok) @@ -345,7 +332,7 @@ class GlobalStateMachine: raise ValueError( f"用户工具集只允许包含第三方工具,以下不合法:{invalid}" ) - saved = await self.postgres_database.upsert_custom_toolset.remote( + saved = await self.postgres_database.upsert_custom_toolset( toolset_id=toolset_id, name=name, tools=list(tools), @@ -359,14 +346,14 @@ class GlobalStateMachine: self._publish_snapshot() return saved - def list_custom_toolsets(self) -> List[Dict[str, Any]]: + async def list_custom_toolsets(self) -> List[Dict[str, Any]]: return list(self._custom_toolsets.values()) - def get_custom_toolset(self, toolset_id: str) -> Optional[Dict[str, Any]]: + async def get_custom_toolset(self, toolset_id: str) -> Optional[Dict[str, Any]]: return self._custom_toolsets.get(toolset_id) async def delete_custom_toolset(self, toolset_id: str) -> bool: - ok = await self.postgres_database.delete_custom_toolset.remote(toolset_id) + ok = await self.postgres_database.delete_custom_toolset(toolset_id) self._custom_toolsets.pop(toolset_id, None) self._global_tool_manager.rebuild_custom_toolsets(self._custom_toolsets) self._publish_snapshot() @@ -374,33 +361,33 @@ class GlobalStateMachine: # ─── Skill ──────────────────────────────────────────────── - def add_skill(self, skill_name: str): + async def add_skill(self, skill_name: str): result = self._global_skill_manager.add_skill(skill_name) self._publish_snapshot() return result - def get_skill_list(self): + async def get_skill_list(self): return self._global_skill_manager.get_skill_list() - def remove_skill(self, skill_name: str): + async def remove_skill(self, skill_name: str): result = self._global_skill_manager.remove_skill(skill_name) self._publish_snapshot() return result # ─── Individual ─────────────────────────────────────────── - def add_individual(self, agent_id: str, config): + async def add_individual(self, agent_id: str, config): result = self._global_individual_manager.add_individual(agent_id, config) self._publish_snapshot() return result - def get_individual(self, agent_id: str): + async def get_individual(self, agent_id: str): return self._global_individual_manager.get_individual(agent_id) - def remove_individual(self, agent_id: str): + async def remove_individual(self, agent_id: str): result = self._global_individual_manager.remove_individual(agent_id) self._publish_snapshot() return result - def list_individuals(self): + async def list_individuals(self): return self._global_individual_manager.list_individuals() diff --git a/kilostar/core/global_state_machine/gsm_snapshot.py b/kilostar/core/global_state_machine/gsm_snapshot.py index ca1eef4..9a2429d 100644 --- a/kilostar/core/global_state_machine/gsm_snapshot.py +++ b/kilostar/core/global_state_machine/gsm_snapshot.py @@ -14,14 +14,14 @@ """GSM 快照对象与客户端拉取工具。 -设计动机:把 GSM actor 内存里的"读路径配置"打包成可放进 Ray Object Store -的不可变快照。读端不再走 actor RPC,而是 ``fetch_snapshot()`` 一次拿到全量 -当前配置,亚毫秒级共享内存读,绕开单 actor 的吞吐瓶颈。 +设计动机:把 GSM 内存里的"读路径配置"打包成一个不可变快照对象。读端不逐字段 +走 GSM 方法,而是 ``fetch_snapshot()`` 一次拿到全量当前配置,进程内直接引用, +配合版本号做进程内缓存失效。 GSM 仍然是 source of truth + 写入串行化器,但读路径解耦: -- 写入 → GSM 内存更新 → ``ray.put(snapshot)`` 拿到新 ObjectRef → 版本号 +1 -- 读取 → ``current_config_ref()`` 拿 (version, ref) → ``ray.get(ref)`` 直读 +- 写入 → GSM 内存更新 → 重建 snapshot 对象 → 版本号 +1 +- 读取 → ``current_version()`` 校验缓存新鲜度 → ``current_config_ref()`` 拿快照 旧的 ``get_provider / get_individual / ...`` 接口保留不动,是低频路径的兜底; 新代码(特别是 skill task 这种高并发热路径)应优先走 ``fetch_snapshot``。 @@ -29,15 +29,9 @@ GSM 仍然是 source of truth + 写入串行化器,但读路径解耦: from __future__ import annotations -import asyncio from dataclasses import dataclass, field from typing import Any, Callable, Dict, List, Optional, Tuple -from kilostar.utils.ray_compat import _STANDALONE - -if not _STANDALONE: - import ray - from kilostar.core.global_state_machine.model_provider.base_provider import Provider from kilostar.utils.logger import get_logger @@ -73,7 +67,6 @@ class GSMSnapshot: _local_cache: Dict[str, Any] = {"version": -1, "snapshot": None} -_cache_lock = asyncio.Lock() async def fetch_snapshot( @@ -83,54 +76,34 @@ async def fetch_snapshot( ) -> GSMSnapshot: """拉取当前 GSM 快照。 - 优先走"版本号检查 + ObjectRef 直读"路径: + 单进程 asyncio 模型下,写入路径原子发布快照对象,读端不需要加锁: - 1. 调 ``gsm.current_version.remote()`` 看本地缓存是否还新(一次轻量 RPC) - 2. 若本地缓存版本号一致,直接返回缓存(亚毫秒,零网络) - 3. 否则调 ``gsm.current_config_ref.remote()`` 拿 ref,``ray.get`` 解出 - 4. 更新本地缓存 - - Args: - use_cache: 默认开启进程内 LRU 缓存(实际是单槽位,持有当前版本); - 测试或诊断场景可关掉强制重拉。 - gsm_actor: 可选传入 GSM actor handle;省略时通过 ``ray_actor_hook`` 获取。 - - Note: - 本函数在 task / actor 进程内多次调用是廉价的;建议每次需要 config 时 - 现取,不要把 snapshot 长期持有跨任务边界(避免 ObjectRef 阻碍回收)。 + 1. 先看本地缓存版本号是否与 ``current_version()`` 一致; + 2. 一致则直接返回缓存对象(少一次全量数据传递); + 3. 不一致则调 ``current_config_ref()`` 拿最新快照并更新缓存。 """ if gsm_actor is None: - from kilostar.utils.ray_hook import ray_actor_hook + from kilostar.utils.actor import get_gsm - gsm_actor = ray_actor_hook("global_state_machine").global_state_machine + gsm_actor = get_gsm() if use_cache: - async with _cache_lock: - try: - latest_version = await gsm_actor.current_version.remote() - except Exception: - latest_version = None + try: + latest_version = await gsm_actor.current_version() + except Exception: + latest_version = None - if ( - latest_version is not None - and _local_cache.get("version") == latest_version - and _local_cache.get("snapshot") is not None - ): - return _local_cache["snapshot"] + if ( + latest_version is not None + and _local_cache["version"] == latest_version + and _local_cache["snapshot"] is not None + ): + return _local_cache["snapshot"] - version, ref_or_snapshot = await gsm_actor.current_config_ref.remote() - if _STANDALONE: - snapshot = ref_or_snapshot - else: - snapshot = ray.get(ref_or_snapshot) - _local_cache["version"] = version - _local_cache["snapshot"] = snapshot - return snapshot - - version, ref_or_snapshot = await gsm_actor.current_config_ref.remote() - if _STANDALONE: - return ref_or_snapshot - return ray.get(ref_or_snapshot) + version, snapshot = await gsm_actor.current_config_ref() + _local_cache["version"] = version + _local_cache["snapshot"] = snapshot + return snapshot def reset_local_cache() -> None: diff --git a/kilostar/core/global_state_machine/individual_manager.py b/kilostar/core/global_state_machine/individual_manager.py index 6de3eaa..10c125d 100644 --- a/kilostar/core/global_state_machine/individual_manager.py +++ b/kilostar/core/global_state_machine/individual_manager.py @@ -32,7 +32,7 @@ class GlobalIndividualManager: """ try: try: - individuals = await postgres.get_all_worker_individual.remote() + individuals = await postgres.get_all_worker_individual() for ind in individuals: agent_id = getattr(ind, "agent_id", None) if agent_id: diff --git a/kilostar/core/global_state_machine/provider_manager.py b/kilostar/core/global_state_machine/provider_manager.py index 889119b..372fec0 100644 --- a/kilostar/core/global_state_machine/provider_manager.py +++ b/kilostar/core/global_state_machine/provider_manager.py @@ -46,7 +46,7 @@ class ProviderManager: async def init_provider_register(self, postgres) -> None: """从 Postgres 读取已存的 Provider 列表,按 provider_title 装入内存注册表。""" - providers = await postgres.get_provider.remote() + providers = await postgres.get_provider() for provider in providers: self.provider_register[provider.provider_title] = provider @@ -90,7 +90,7 @@ class ProviderManager: provider: Provider = await provider_class.create_provider(provider_args) provider.provider_owner = provider_owner self.provider_register[provider_title] = provider - await postgres_database.add_provider_db.remote( + await postgres_database.add_provider_db( provider_id=str(ulid.ULID()), provider_title=provider.provider_title, provider_url=provider.provider_url, @@ -125,7 +125,7 @@ class ProviderManager: async def delete_provider(self, provider_title: str, postgres_database) -> None: """从内存注册表 + Postgres 中一并删除指定 Provider;不存在时静默返回。""" if provider_title in self.provider_register: - await postgres_database.delete_provider_by_title.remote( + await postgres_database.delete_provider_by_title( provider_title=provider_title ) del self.provider_register[provider_title] diff --git a/kilostar/core/global_workflow_manager/global_workflow_manager.py b/kilostar/core/global_workflow_manager/global_workflow_manager.py index bddf9c6..e033e48 100644 --- a/kilostar/core/global_workflow_manager/global_workflow_manager.py +++ b/kilostar/core/global_workflow_manager/global_workflow_manager.py @@ -1,7 +1,5 @@ import asyncio from typing import Dict -from kilostar.utils.ray_compat import actor_class -from kilostar.utils.ray_hook import ray_actor_hook from kilostar.utils.logger import get_logger @@ -11,7 +9,6 @@ class TraceQueues: self.receive: asyncio.Queue[str] = asyncio.Queue() -@actor_class class GlobalWorkflowManager: def __init__(self): self._traces: Dict[str, TraceQueues] = {} diff --git a/kilostar/core/individual/consciousness_node/consciousness_node.py b/kilostar/core/individual/consciousness_node/consciousness_node.py index 0d317b6..2213e1b 100644 --- a/kilostar/core/individual/consciousness_node/consciousness_node.py +++ b/kilostar/core/individual/consciousness_node/consciousness_node.py @@ -14,7 +14,6 @@ from typing import Union, overload -from kilostar.utils.ray_compat import actor_class from kilostar.core.individual.consciousness_node.template import ( ConsciousnessNodeDeps, ForregulatoryNode, @@ -28,11 +27,9 @@ from pydantic_ai import Agent, RunContext from kilostar.core.global_state_machine.global_state_machine import GlobalStateMachine from kilostar.core.global_state_machine.model_provider.base_provider import Provider from kilostar.adapter.model_adapter.agent_factory import AgentFactory -from kilostar.utils.ray_hook import ray_actor_hook from kilostar.utils.prompts import agent_prompt -@actor_class class ConsciousnessNode: def __init__(self) -> None: from kilostar.utils.logger import get_logger @@ -134,8 +131,10 @@ class ConsciousnessNode: f"ConsciousnessNode: 开始为 trace_id {trace_id} 设计工作流。原始命令:{command}" ) # 获取可用技能 (示例) - postgres_database = ray_actor_hook("postgres_database").postgres_database - skills_entities = await postgres_database.get_all_worker_individual.remote() + from kilostar.utils.actor import get_postgres + + postgres_database = get_postgres() + skills_entities = await postgres_database.get_all_worker_individual() available_skills = [] if skills_entities: for skill in skills_entities: @@ -151,10 +150,10 @@ class ConsciousnessNode: original_command=command, available_skills=available_skills ) - global_workflow_manager = ray_actor_hook( - "global_workflow_manager" - ).global_workflow_manager - await global_workflow_manager.put_pending.remote( + from kilostar.utils.actor import get_gwm + + global_workflow_manager = get_gwm() + await global_workflow_manager.put_pending( trace_id, "正在为您构建并规划工作流任务节点,请稍候..." ) @@ -164,20 +163,21 @@ class ConsciousnessNode: workflow = result.workflow workflow.trace_id = trace_id - await global_workflow_manager.put_pending.remote( + await global_workflow_manager.put_pending( trace_id, "工作流构建完成,即将开始执行!" ) - # 直接以 ray task 形式 fire workflow,不再经过 WorkflowRunningEngine 这层中转: - # workflow 是一次性、有头有尾的执行,task 语义比常驻 actor 更贴。 + # 以后台任务形式 fire workflow:workflow 是一次性、有头有尾的 + # 执行,跑完即结束;持久化交给 PostgresStatePersistence 保证可恢复。 + from kilostar.utils.actor import spawn_background from kilostar.core.work.workflow.workflow_engine import run_workflow_task - run_workflow_task.remote(workflow.model_dump(), trace_id) + spawn_background(run_workflow_task(workflow.model_dump(), trace_id)) else: - await global_workflow_manager.put_pending.remote( + await global_workflow_manager.put_pending( trace_id, "很抱歉,工作流生成失败。" ) - await postgres_database.update_workflow_status.remote(trace_id, "failed") + await postgres_database.update_workflow_status(trace_id, "failed") async def working( self, diff --git a/kilostar/core/individual/consciousness_node/template.py b/kilostar/core/individual/consciousness_node/template.py index 52d2ff5..b492d0b 100644 --- a/kilostar/core/individual/consciousness_node/template.py +++ b/kilostar/core/individual/consciousness_node/template.py @@ -14,24 +14,23 @@ from kilostar.core.work.workflow.workflow import KiloStarWorkflow, WorkflowStep -from kilostar.utils.agent_model import ResponseModel, DepsModel, RequestModel -from pydantic import Field +from pydantic import BaseModel, Field from typing import Optional, List # 意识节点回复类 -class ConsciousnessNodeResponse(ResponseModel): +class ConsciousnessNodeResponse(BaseModel): """Consciousness response model,是意识节点所有回复类型的父类""" pass -class ConsciousnessNodeDeps(DepsModel): +class ConsciousnessNodeDeps(BaseModel): """ConsciousnessNode 在 pydantic-ai Agent 中使用的依赖:原始指令、当前指令以及可用 Skill 列表。""" original_command: str command: str available_skills: Optional[List[dict]] = None locale: str = "zh" -class ConsciousnessNodeInput(RequestModel): +class ConsciousnessNodeInput(BaseModel): """ConsciousnessNode 各类入参的共同基类,仅用于打 schema 标签。""" pass diff --git a/kilostar/core/individual/control_node/__init__.py b/kilostar/core/individual/control_node/__init__.py deleted file mode 100644 index 9948e4e..0000000 --- a/kilostar/core/individual/control_node/__init__.py +++ /dev/null @@ -1,17 +0,0 @@ -# Copyright 2026 zhaoxi826 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from .control_node import ControlNode - -__all__ = ["ControlNode"] diff --git a/kilostar/core/individual/control_node/control_node.py b/kilostar/core/individual/control_node/control_node.py deleted file mode 100644 index 802c31e..0000000 --- a/kilostar/core/individual/control_node/control_node.py +++ /dev/null @@ -1,135 +0,0 @@ -# Copyright 2026 zhaoxi826 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from pydantic_ai import Agent, RunContext -from kilostar.utils.ray_compat import actor_class -from kilostar.core.global_state_machine.global_state_machine import GlobalStateMachine -from kilostar.core.global_state_machine.model_provider.base_provider import Provider -from kilostar.adapter.model_adapter.agent_factory import AgentFactory -from kilostar.core.individual.control_node.template import ( - ForWorkflow, - ForWorkflowInput, - ControlNodeDeps, -) -from kilostar.utils.prompts import agent_prompt - - -@actor_class -class ControlNode: - """ControlNode(控制节点):**已废弃**——名字保留给未来的远程探针/系统控制节点。 - - 历史:早期设计里它是工作流的"单步执行 actor",但 workflow_engine 的 Dispatch - 最终只识别 ``consciousness_node`` 和 ``skill_individual``,本类从未真正被调用过。 - 保留目录与类壳子,避免改名带来的 git 历史断层;**不要新增对它的依赖**。 - 待远程探针/监控流子项目启动时,本目录将被重写为远程机器控制节点。 - """ - - def __init__(self): - from kilostar.utils.logger import get_logger - - self.logger = get_logger("control_node") - self.agent: Agent | None = None - self._model_settings: dict = {} - - async def create_agent( - self, - global_state_machine: GlobalStateMachine, - provider_title: str, - model_id: str, - tools_list: list[str] = None, - toolsets=None, - locale: str | None = None, - custom_system_prompt: str | None = None, - ) -> None: - """ - create_agent方法,将agent对象装配到Control的属性内 - 该方法通过provider_title从global_state_machine中获取provider对象,然后从provider对象中取出供应商形象,装配为pydantic_ai的 - Agent实例, - 并挂载到self.agent属性 - Args: - global_state_machine: 全局状态机 - provider_title: 供应商名 - model_id: 模型id - locale: 语言代码(zh/en),控制system prompt语言 - custom_system_prompt: 管理员自定义追加提示词(可选) - - Returns: - 无返回 - """ - system_prompt: str = agent_prompt("control_node", locale=locale, custom_system_prompt=custom_system_prompt) - output_type = ForWorkflow - from kilostar.utils.get_tool import load_tools_from_list - from kilostar.core.global_state_machine.gsm_snapshot import fetch_snapshot - - snapshot = await fetch_snapshot(gsm_actor=global_state_machine) - provider: Provider = snapshot.providers.get(provider_title) - if provider is None: - from kilostar.utils.i18n import t - raise ValueError(t("provider_not_registered", locale=locale, provider_title=provider_title)) - agent_factory = AgentFactory() - - callables = load_tools_from_list(tools_list) - self.agent = agent_factory.create_agent( - provider=provider, - model_id=model_id, - output_type=output_type, - system_prompt=system_prompt, - deps_type=ControlNodeDeps, - agent_name="control_node", - tools=callables, - toolsets=toolsets, - ) - self._model_settings = AgentFactory.resolve_model_settings(provider, model_id) - - @self.agent.system_prompt - async def dynamic_prompt(ctx: RunContext[ControlNodeDeps]): - """运行期动态拼接 system prompt:把当前 workflow_step 的关键字段塞进去。""" - prompt = system_prompt + "\n\n" - prompt += ( - f"=== 当前任务步骤上下文 ===\n" - f"- 步骤名称 (Name): {ctx.deps.workflow_step.name}\n" - f"- 步骤目标/描述 (Description): {ctx.deps.workflow_step.desc}\n" - f"- 前置输入(input): {ctx.deps.workflow_step.inputs}\n" - ) - return prompt - - async def working(self, payload: ForWorkflowInput) -> str: - """对外入口:执行一次步骤,吞掉异常并返回 ``None`` 以避免拖垮上游 Workflow。""" - try: - result: ForWorkflow = await self._run(payload) - return result - except Exception: - self.logger.exception("ControlNode在执行working时发生严重错误") - return None - - async def _run(self, payload: ForWorkflowInput) -> ForWorkflow: - """实际执行步骤:组装 ``ControlNodeDeps``、调用 Agent,最终把 ``ForWorkflow`` 输出取出。""" - try: - self.agent.retries = 3 - deps = ControlNodeDeps(workflow_step=payload.workflow_step) - self.logger.debug( - f"ControlNode: 开始执行工作流节点 [{payload.workflow_step.name}] (原生重试开启)" - ) - - result = await self.agent.run( - f"请根据提供的 workflow_step 上下文,执行此步骤并输出结果。\n详细指令或附加数据:{payload.workflow_step.model_dump_json()}", - deps=deps, - model_settings=self._model_settings or None, - ) - return result.output - except Exception as e: - self.logger.exception( - f"ControlNode 在执行步骤 [{payload.workflow_step.name}] 时最终失败: {str(e)}" - ) - raise RuntimeError(f"ControlNode 执行步骤失败: {str(e)}") from e diff --git a/kilostar/core/individual/control_node/template.py b/kilostar/core/individual/control_node/template.py deleted file mode 100644 index c4ed5a3..0000000 --- a/kilostar/core/individual/control_node/template.py +++ /dev/null @@ -1,51 +0,0 @@ -# Copyright 2026 zhaoxi826 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -from pydantic import Field -from kilostar.core.work.workflow.workflow import WorkflowStep -from kilostar.utils.agent_model import ResponseModel, RequestModel, DepsModel - - -class ControlNodeResponse(ResponseModel): - """控制节点回复的基类。""" - - pass - - -class ControlNodeInput(RequestModel): - """控制节点输入的基类,承载一次调度所需的入参。""" - - pass - - -class ControlNodeDeps(DepsModel): - """控制节点运行期依赖,注入到 pydantic-ai Agent 的 RunContext。""" - - workflow_step: WorkflowStep - # In the future, this can be dynamically populated with tools specific to the current task execution - - -class ForWorkflow(ControlNodeResponse): - """控制节点执行单个工作流步骤的输出模型。""" - - output: str = Field( - ..., description="控制节点执行特定工作流步骤的结果。包含执行细节和输出数据。" - ) - - -class ForWorkflowInput(ControlNodeInput): - """控制节点针对工作流步骤的输入模型。""" - - workflow_step: WorkflowStep diff --git a/kilostar/core/individual/growth_node/__init__.py b/kilostar/core/individual/growth_node/__init__.py deleted file mode 100644 index 5fa7362..0000000 --- a/kilostar/core/individual/growth_node/__init__.py +++ /dev/null @@ -1,14 +0,0 @@ -# Copyright 2026 zhaoxi826 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - diff --git a/kilostar/core/individual/growth_node/growth_node.py b/kilostar/core/individual/growth_node/growth_node.py deleted file mode 100644 index 5fa7362..0000000 --- a/kilostar/core/individual/growth_node/growth_node.py +++ /dev/null @@ -1,14 +0,0 @@ -# Copyright 2026 zhaoxi826 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - diff --git a/kilostar/core/individual/regulatory_node/regulatory_node.py b/kilostar/core/individual/regulatory_node/regulatory_node.py index 8bfe44c..1588500 100644 --- a/kilostar/core/individual/regulatory_node/regulatory_node.py +++ b/kilostar/core/individual/regulatory_node/regulatory_node.py @@ -15,7 +15,6 @@ import asyncio import datetime from typing import Union -from kilostar.utils.ray_compat import actor_class from kilostar.adapter.model_adapter.agent_factory import AgentFactory from kilostar.core.global_state_machine.global_state_machine import GlobalStateMachine from kilostar.core.global_state_machine.model_provider import Provider @@ -28,7 +27,6 @@ from pydantic_ai import RunContext, Agent from kilostar.utils.prompts import agent_prompt -@actor_class class RegulatoryNode: """RegulatoryNode(监管节点):用户请求的直接对话节点。 @@ -234,12 +232,12 @@ class RegulatoryNode: return try: import uuid - from kilostar.utils.ray_hook import ray_actor_hook + from kilostar.utils.actor import get_actor - postgres_database = ray_actor_hook("postgres_database").postgres_database + postgres_database = get_actor("postgres_database") task_id = uuid.uuid4().hex chat_id = payload.platform_id if payload.platform == "client" else None - await postgres_database.create_task.remote( + await postgres_database.create_task( task_id=task_id, user_id=payload.user_name, command=payload.message, diff --git a/kilostar/core/individual/regulatory_node/template.py b/kilostar/core/individual/regulatory_node/template.py index 907b461..8fd2fbf 100644 --- a/kilostar/core/individual/regulatory_node/template.py +++ b/kilostar/core/individual/regulatory_node/template.py @@ -13,12 +13,10 @@ # limitations under the License. from typing import Literal, Optional -from pydantic import Field - -from kilostar.utils.agent_model import ResponseModel, DepsModel, RequestModel +from pydantic import BaseModel, Field -class RegulatoryNodeResponse(ResponseModel): +class RegulatoryNodeResponse(BaseModel): """ RegulatoryNodeResponse类 一切regulatory_node回复的父类 @@ -26,14 +24,14 @@ class RegulatoryNodeResponse(ResponseModel): pass -class RegulatoryNodeRequest(RequestModel): +class RegulatoryNodeRequest(BaseModel): """ RegulatoryNodeRequest类 向regulatory请求的父类 """ pass -class RegulatoryNodeDeps(DepsModel): +class RegulatoryNodeDeps(BaseModel): """ RegulatoryNodeDeps类 regulatory_node的依赖模型 @@ -44,7 +42,7 @@ class RegulatoryNodeDeps(DepsModel): retry_count: int = 0 error_history: str = "" -class MessageRequest(RequestModel): +class MessageRequest(BaseModel): """ MessageRequest类 任何消息渠道向regulatory_node发送消息请求的模型 diff --git a/kilostar/core/postgres_database/postgres.py b/kilostar/core/postgres_database/postgres.py index c95a381..b347e48 100644 --- a/kilostar/core/postgres_database/postgres.py +++ b/kilostar/core/postgres_database/postgres.py @@ -15,7 +15,6 @@ import os import asyncio -from kilostar.utils.ray_compat import actor_class from kilostar.utils.settings import get_settings from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession from sqlalchemy.orm import sessionmaker @@ -62,9 +61,8 @@ from .module.org_task import OrgTaskDatabase from .module.task import TaskDatabase -@actor_class class PostgresDatabase: - """以 Ray Actor 形式暴露的统一数据库门面。 + """经 actor 运行时(``get_actor("postgres_database")``)暴露的统一数据库门面。 内部组合了 Auth / Provider / Individual / SystemNode / Workflow / ChatHistory 六个子库,所有方法在调用前都会等待 ``ready_event``,确保 ``init_db`` 完成后 diff --git a/kilostar/core/work/workflow/graph_persistence.py b/kilostar/core/work/workflow/graph_persistence.py index 6df5edb..e7a9a30 100644 --- a/kilostar/core/work/workflow/graph_persistence.py +++ b/kilostar/core/work/workflow/graph_persistence.py @@ -169,16 +169,16 @@ def _from_json_bytes(blob: bytes) -> Any: def build_postgres_persistence(trace_id: str) -> PostgresStatePersistence: - """生产环境构造 PostgresStatePersistence:从 ray_actor_hook 取 postgres handle。""" - from kilostar.utils.ray_hook import ray_actor_hook + """生产环境构造 PostgresStatePersistence:取进程内 postgres 服务。""" + from kilostar.utils.actor import get_postgres - postgres_database = ray_actor_hook("postgres_database").postgres_database + postgres_database = get_postgres() async def _write(tid: str, history: Any) -> None: - await postgres_database.upsert_workflow_graph_state.remote(tid, history) + await postgres_database.upsert_workflow_graph_state(tid, history) async def _read(tid: str) -> Optional[Any]: - record = await postgres_database.get_workflow_graph_state.remote(tid) + record = await postgres_database.get_workflow_graph_state(tid) if record is None: return None # ORM 模型 / dict / list 都兼容 diff --git a/kilostar/core/work/workflow/workflow_engine.py b/kilostar/core/work/workflow/workflow_engine.py index e26b8c3..0b765ba 100644 --- a/kilostar/core/work/workflow/workflow_engine.py +++ b/kilostar/core/work/workflow/workflow_engine.py @@ -36,7 +36,6 @@ import datetime from dataclasses import dataclass from typing import Any, Awaitable, Callable, Dict, List, Optional -from kilostar.utils.ray_compat import remote_task, _STANDALONE from pydantic import BaseModel, Field from pydantic_graph import BaseNode, End, Graph, GraphRunContext from pydantic_graph.persistence import BaseStatePersistence @@ -74,7 +73,7 @@ StepExecutor = Callable[ class WorkflowDeps: """节点运行期依赖:所有外部 IO 都从这里走,便于测试 mock。 - 每个字段都是一个 awaitable,签名贴近原 ``.remote()`` 调用。生产路径由 + 每个字段都是一个 awaitable,签名贴近底层 actor 方法。生产路径由 ``build_default_deps`` 现场组装真实 actor handle 包装;单测可以传任意 ``AsyncMock``。 @@ -419,8 +418,8 @@ async def _default_skill_executor( async def _default_consciousness_executor( step_data: Dict[str, Any], state: WorkflowGraphState ) -> tuple[str, bool]: - """生产环境的 consciousness 派发器:远程调用 ConsciousnessNode.working。""" - from kilostar.utils.ray_hook import ray_actor_hook + """生产环境的 consciousness 派发器:调用 ConsciousnessNode.working。""" + from kilostar.utils.actor import get_consciousness from kilostar.core.individual.consciousness_node.template import ( ForWorkflow, ForWorkflowInput, @@ -428,12 +427,12 @@ async def _default_consciousness_executor( ) from kilostar.core.work.workflow.workflow import WorkflowStep - consciousness_node = ray_actor_hook("consciousness_node").consciousness_node + consciousness_node = get_consciousness() payload = ForWorkflowInput( workflow_step=WorkflowStep.model_validate(step_data), original_command=state.original_command, ) - result = await consciousness_node.working.remote(payload) + result = await consciousness_node.working(payload) if isinstance(result, ForWorkflow): return result.output, True if isinstance(result, ForregulatoryNode): @@ -444,32 +443,30 @@ async def _default_consciousness_executor( def build_default_deps() -> WorkflowDeps: - """生产环境构造 ``WorkflowDeps``:把 ray actor handle 包装成 awaitable。 + """生产环境构造 ``WorkflowDeps``:把进程内服务包装成 awaitable。 抽出来是为了让 ``run_workflow_task`` 入口和测试入口共享同一套包装逻辑。 """ - from kilostar.utils.ray_hook import ray_actor_hook + from kilostar.utils.actor import get_postgres, get_gwm - postgres_database = ray_actor_hook("postgres_database").postgres_database - global_workflow_manager = ray_actor_hook( - "global_workflow_manager" - ).global_workflow_manager + postgres_database = get_postgres() + global_workflow_manager = get_gwm() async def _upsert_workflow_context(trace_id: str, **kwargs: Any) -> Any: - return await postgres_database.upsert_workflow_context.remote( + return await postgres_database.upsert_workflow_context( trace_id, **kwargs ) async def _update_workflow_status(trace_id: str, status: str) -> Any: - return await postgres_database.update_workflow_status.remote( + return await postgres_database.update_workflow_status( trace_id, status ) async def _put_pending(trace_id: str, message: str) -> Any: - return await global_workflow_manager.put_pending.remote(trace_id, message) + return await global_workflow_manager.put_pending(trace_id, message) async def _get_received(trace_id: str) -> str: - return await global_workflow_manager.get_received.remote(trace_id) + return await global_workflow_manager.get_received(trace_id) return WorkflowDeps( upsert_workflow_context=_upsert_workflow_context, @@ -547,11 +544,13 @@ async def resume_workflow_graph( return final_output -@remote_task -def run_workflow_task( +async def run_workflow_task( workflow_data: dict, trace_id: str, resume_only: bool = False -): - """workflow 的 ray task 入口:一次性执行,跑完即销毁。 +) -> None: + """workflow 的后台任务入口:一次性执行,跑完即结束。 + + 调用方用 ``asyncio.create_task(run_workflow_task(...))`` 以 fire-and-forget + 方式在当前事件循环里后台执行。 生产路径下持久化交给 ``PostgresStatePersistence`` —— 即便进程崩溃,再 fire 一次相同 ``trace_id`` 的任务(或调 ``/workflow/{trace_id}/resume``)即可 @@ -564,8 +563,7 @@ def run_workflow_task( 否则会拿着空 ``workflow_data`` 空跑一个 ``work_link=[]`` 的 workflow 并误判 为 COMPLETED(静默 bug)。 - ray task 是新进程,contextvars 不会从 caller 传过来,所以入口先 bind 一次 - ``trace_id``,让节点内的日志自动带上它。 + 入口先 bind 一次 ``trace_id``,让节点内的日志自动带上它。 """ from kilostar.utils.request_context import trace_id_scope from kilostar.core.work.workflow.graph_persistence import ( @@ -575,35 +573,29 @@ def run_workflow_task( _logger = get_logger("workflow_task") - async def _entry() -> None: - with trace_id_scope(trace_id): - persistence = build_postgres_persistence(trace_id) - persistence.set_graph_types(workflow_graph) + with trace_id_scope(trace_id): + persistence = build_postgres_persistence(trace_id) + persistence.set_graph_types(workflow_graph) + recovered = False + try: + recovered = await persistence.hydrate() + except Exception as e: + if resume_only: + _logger.error(f"resume 失败:无法 hydrate 图持久化记录: {e}") + raise recovered = False - try: - recovered = await persistence.hydrate() - except Exception as e: - if resume_only: - _logger.error(f"resume 失败:无法 hydrate 图持久化记录: {e}") - raise - recovered = False - if resume_only and not recovered: - msg = ( - f"resume 失败:trace {trace_id} 没有可恢复的图持久化记录," - "拒绝以全新模式空跑" - ) - _logger.error(msg) - raise RuntimeError(msg) + if resume_only and not recovered: + msg = ( + f"resume 失败:trace {trace_id} 没有可恢复的图持久化记录," + "拒绝以全新模式空跑" + ) + _logger.error(msg) + raise RuntimeError(msg) - if recovered: - await resume_workflow_graph(trace_id, persistence=persistence) - else: - await run_workflow_graph( - workflow_data, trace_id, persistence=persistence - ) - - if _STANDALONE: - return _entry() - else: - asyncio.run(_entry()) + if recovered: + await resume_workflow_graph(trace_id, persistence=persistence) + else: + await run_workflow_graph( + workflow_data, trace_id, persistence=persistence + ) diff --git a/kilostar/plugin_runtime/base_organization.py b/kilostar/plugin_runtime/base_organization.py index 326f774..7f72d0b 100644 --- a/kilostar/plugin_runtime/base_organization.py +++ b/kilostar/plugin_runtime/base_organization.py @@ -1,7 +1,7 @@ """BaseOrganization:重型插件基类。 设计要点: -- 单机模式 = 普通 Python 对象,分布式 = ray actor(``@actor_class`` 装饰子类) +- 进程内普通 Python 对象,经 ``register_actor`` 登记后由 ``get_actor`` 寻址 - 内置 ``asyncio.Queue`` 输入队列 + 任务表 - 对外两条通道:``dispatch`` (阻塞) / ``submit`` (射后不管),底层都汇集到 ``_run_task`` - 子类只需覆写 ``setup`` / ``react`` 两个钩子;零代码插件由 ``agents.json`` 声明驱动 @@ -335,10 +335,10 @@ class BaseOrganization: async def _persist_task(self, ts: TaskState) -> None: """把任务状态写到 PG。失败不阻塞执行。""" try: - from kilostar.utils.ray_hook import ray_actor_hook + from kilostar.utils.actor import get_actor - pg = ray_actor_hook("postgres_database").postgres_database - await pg.upsert_org_task.remote( + pg = get_actor("postgres_database") + await pg.upsert_org_task( task_id=ts.task_id, org_name=ts.org_name, trace_id=ts.trace_id, @@ -354,10 +354,10 @@ class BaseOrganization: async def _persist_event(self, ts: TaskState, ev: OrgEvent) -> None: try: - from kilostar.utils.ray_hook import ray_actor_hook + from kilostar.utils.actor import get_actor - pg = ray_actor_hook("postgres_database").postgres_database - await pg.append_org_task_event.remote( + pg = get_actor("postgres_database") + await pg.append_org_task_event( task_id=ts.task_id, event=ev.to_dict() ) except Exception: @@ -510,10 +510,10 @@ class BaseOrganization: if adef.model and adef.model.provider_title and adef.model.model_id: return adef.model.provider_title, adef.model.model_id try: - from kilostar.utils.ray_hook import ray_actor_hook + from kilostar.utils.actor import get_actor - pg = ray_actor_hook("postgres_database").postgres_database - row = await pg.find_plugin_slot.remote(self.name, adef.name) + pg = get_actor("postgres_database") + row = await pg.find_plugin_slot(self.name, adef.name) if row is None: return "", "" return getattr(row, "provider_title", "") or "", getattr(row, "model_id", "") or "" diff --git a/kilostar/plugin_runtime/plugin_manager.py b/kilostar/plugin_runtime/plugin_manager.py index 95b2cac..dbc79a9 100644 --- a/kilostar/plugin_runtime/plugin_manager.py +++ b/kilostar/plugin_runtime/plugin_manager.py @@ -19,18 +19,16 @@ from kilostar.plugin_runtime.loader import ( from kilostar.plugin_runtime.manifest import OrgManifest from kilostar.plugin_runtime.tool_bridge import make_dispatch_tool from kilostar.utils.logger import get_logger -from kilostar.utils.ray_compat import _STANDALONE, actor_class -from kilostar.utils.ray_hook import register_standalone +from kilostar.utils.actor import register_actor from kilostar.utils.settings import get_plugin_data_dir, get_plugin_dir logger = get_logger("plugin_manager") -@actor_class class GlobalPluginManager: - """单机模式下是对象,分布式下是 ray actor。 + """插件运行时管理器(进程内单例)。 - 每个 loaded 组织保存其 manifest 和 actor handle(standalone=proxy,dist=ray handle)。 + 每个 loaded 组织保存其 manifest 和实例,并注册到进程内 actor 运行时。 """ def __init__(self): @@ -64,13 +62,18 @@ class GlobalPluginManager: org_info = self._orgs.pop(name, None) if org_info is None: return {"name": name, "status": "not_found"} - # shutdown actor + # shutdown 实例并从 actor 运行时注销 try: handle = org_info.get("handle") if handle is not None: - await handle.shutdown.remote() + await handle.shutdown() except Exception as e: logger.warning(f"shutdown org_{name} failed: {e}") + actor_name = org_info.get("actor_name") + if actor_name: + from kilostar.utils.actor import unregister_actor + + unregister_actor(actor_name) # 移除 dispatch tool self._dispatch_tools.pop(f"dispatch_to_{name}", None) logger.info(f"uninstalled plugin: {name}") @@ -136,15 +139,9 @@ class GlobalPluginManager: await instance.setup() - # 注册到 ray_actor_hook 命名空间 + # 注册到进程内 actor 运行时,业务侧通过 get_actor(actor_name) 取用 actor_name = manifest.actor_name - if _STANDALONE: - register_standalone(actor_name, instance) - else: - # 分布式模式下,这里需要把 instance 包装成 ray actor - # 第一版走 standalone 逻辑(两种模式统一 register 到本进程) - # 真正分布式隔离等后续做 - register_standalone(actor_name, instance) + register_actor(actor_name, instance) # 生成 dispatch tool tool = make_dispatch_tool(name, manifest.display_name, manifest.description) @@ -166,15 +163,15 @@ class GlobalPluginManager: DB 不可用时静默跳过(standalone 启动早期 / 单测场景)。 """ try: - from kilostar.utils.ray_hook import ray_actor_hook + from kilostar.utils.actor import get_actor - pg = ray_actor_hook("postgres_database").postgres_database + pg = get_actor("postgres_database") for adef in agents_dict.get("agents", []): slot_name = adef.get("name") if not slot_name: continue description = adef.get("role") or adef.get("system_prompt") or slot_name - await pg.upsert_plugin_slot.remote( + await pg.upsert_plugin_slot( plugin_name=name, slot_name=slot_name, description=description, @@ -185,10 +182,10 @@ class GlobalPluginManager: async def cleanup_orphan_plugin_slots(self) -> None: """启动期兜底:DB 中存在但目录已不在的 plugin_owned slot 全部清掉。""" try: - from kilostar.utils.ray_hook import ray_actor_hook + from kilostar.utils.actor import get_actor - pg = ray_actor_hook("postgres_database").postgres_database - recorded: List[str] = await pg.list_plugin_owned_names.remote() or [] + pg = get_actor("postgres_database") + recorded: List[str] = await pg.list_plugin_owned_names() or [] except Exception as e: logger.debug(f"cleanup_orphan_plugin_slots skipped: {e}") return @@ -198,7 +195,7 @@ class GlobalPluginManager: for plugin_name in recorded: if plugin_name not in present: try: - n = await pg.delete_plugin_slots.remote(plugin_name) + n = await pg.delete_plugin_slots(plugin_name) logger.info(f"cleaned {n} orphan slots for missing plugin {plugin_name!r}") except Exception as e: logger.warning(f"failed to clean orphan slots for {plugin_name}: {e}") diff --git a/kilostar/plugin_runtime/tool_bridge.py b/kilostar/plugin_runtime/tool_bridge.py index 4b45e37..51421a9 100644 --- a/kilostar/plugin_runtime/tool_bridge.py +++ b/kilostar/plugin_runtime/tool_bridge.py @@ -6,7 +6,7 @@ RegulatoryNode/ConsciousnessNode 通过这个工具向部门派单,等待部 from __future__ import annotations -from typing import Callable, Dict +from typing import Callable def make_dispatch_tool(org_name: str, display_name: str, description: str) -> Callable: @@ -19,12 +19,10 @@ def make_dispatch_tool(org_name: str, display_name: str, description: str) -> Ca desc_text = description or f"把任务派给{display_name or org_name}部门,由部门内部多 agent 协作完成。" async def _impl(task_description: str) -> str: - from kilostar.utils.ray_hook import ray_actor_hook + from kilostar.utils.actor import get_actor - actor_name = f"org_{org_name}" - actor = ray_actor_hook(actor_name) - target = getattr(actor, actor_name) - result = await target.dispatch.remote(task_description, {}) + target = get_actor(f"org_{org_name}") + result = await target.dispatch(task_description, {}) if result.get("status") == "completed": return str(result.get("result") or "") return f"[{org_name} 任务失败] {result.get('error') or 'unknown'}" @@ -38,15 +36,3 @@ def make_dispatch_tool(org_name: str, display_name: str, description: str) -> Ca " 部门交付的结果文本,失败时返回错误说明。\n" ) return _impl - - -def collect_dispatch_tools(org_specs: Dict[str, Dict[str, str]]) -> Dict[str, Callable]: - """根据 ``{org_name: {"display_name": ..., "description": ...}}`` 批量生成。""" - return { - f"dispatch_to_{name}": make_dispatch_tool( - name, - spec.get("display_name", ""), - spec.get("description", ""), - ) - for name, spec in org_specs.items() - } diff --git a/kilostar/utils/access.py b/kilostar/utils/access.py index 33ca3a6..d79cafc 100644 --- a/kilostar/utils/access.py +++ b/kilostar/utils/access.py @@ -191,11 +191,11 @@ async def get_authority(user_id: str) -> "UserAuthority": """通过 PostgresDatabase Actor 查出指定用户的 ``UserAuthority``;用户不存在时抛 401。""" from kilostar.utils.error import UserNotExistError from kilostar.utils.i18n import t - from kilostar.utils.ray_hook import ray_actor_hook + from kilostar.utils.actor import get_actor - postgres_database = ray_actor_hook("postgres_database").postgres_database + postgres_database = get_actor("postgres_database") try: - user_authority = await postgres_database.get_user_authority.remote( + user_authority = await postgres_database.get_user_authority( user_id=user_id ) return user_authority diff --git a/kilostar/utils/actor.py b/kilostar/utils/actor.py new file mode 100644 index 0000000..4c0d9be --- /dev/null +++ b/kilostar/utils/actor.py @@ -0,0 +1,180 @@ +# Copyright 2026 zhaoxi826 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""进程内 actor 运行时(位置透明的组件寻址层)。 + +设计目标:**一处编写,处处运行**。业务代码永远通过 ``get_actor(name)`` 拿到一个 +``ActorHandle``,再 ``await handle.method(args)`` 调用,**不关心对端在哪**。 + +今天所有核心组件(PostgresDatabase / GlobalStateMachine / ...)都是本进程内的 +普通异步对象,``get_actor`` 返回 ``LocalActorHandle``,直接转发到本地实例。 + +未来若要把某一层(如跑 vLLM 的重型 worker)拆到独立进程/机器上,只需新增一个 +``RemoteActorHandle``(gRPC / 子进程 / 消息队列后端)并在注册时选择它——业务侧的 +调用代码一行都不用改。这层接缝就是当初 Ray ``.remote()`` 提供的价值,这里用零依赖 +的方式把它保留下来。 +""" + +from __future__ import annotations + +import asyncio +from typing import Any, Awaitable, Dict, Optional, Set + +_registry: Dict[str, Any] = {} +_handle_cache: Dict[str, "ActorHandle"] = {} + +# 已 spawn 的后台任务:持强引用,防止事件循环只持弱引用导致任务被 GC 掉 +_background_tasks: Set["asyncio.Task[Any]"] = set() + + +class _BoundMethod: + """包装目标对象的一个方法,使调用统一返回 awaitable。 + + - 目标方法是 ``async def`` → 直接返回其协程 + - 目标方法是普通函数 → 把返回值包成一个已完成的 awaitable + 这样调用方永远可以 ``await handle.method(...)``,无需关心同步/异步、本地/远程。 + """ + + __slots__ = ("_fn",) + + def __init__(self, fn: Any) -> None: + self._fn = fn + + def __call__(self, *args: Any, **kwargs: Any) -> Any: + result = self._fn(*args, **kwargs) + if asyncio.iscoroutine(result): + return result + + async def _wrap() -> Any: + return result + + return _wrap() + + +class ActorHandle: + """actor 句柄基类:定义位置透明的调用接口。 + + 未来的 ``RemoteActorHandle`` 继承本类,改写方法解析为跨进程 RPC 即可。 + """ + + +class LocalActorHandle(ActorHandle): + """本进程句柄:方法调用直接转发到已注册的实例。""" + + __slots__ = ("_impl",) + + def __init__(self, impl: Any) -> None: + object.__setattr__(self, "_impl", impl) + + def __getattr__(self, name: str) -> Any: + attr = getattr(object.__getattribute__(self, "_impl"), name) + if callable(attr): + return _BoundMethod(attr) + return attr + + +def register_actor(name: str, instance: Any) -> None: + """注册一个 actor 实例。重复注册会覆盖(支持插件热重载)。""" + _registry[name] = instance + _handle_cache.pop(name, None) + + +def unregister_actor(name: str) -> None: + """注销一个 actor(插件卸载用);不存在时静默。""" + _registry.pop(name, None) + _handle_cache.pop(name, None) + + +def get_actor(name: str) -> ActorHandle: + """按名字取 actor 句柄;未注册时抛 KeyError。 + + 动态名(如 ``f"org_{plugin}"``)也走这个统一入口。 + """ + if name not in _registry: + raise KeyError(f"actor {name!r} 未注册") + handle = _handle_cache.get(name) + if handle is None: + handle = LocalActorHandle(_registry[name]) + _handle_cache[name] = handle + return handle + + +def try_get_actor(name: str) -> Optional[ActorHandle]: + """按名字取 actor 句柄;未注册时返回 None(可选依赖用)。""" + try: + return get_actor(name) + except KeyError: + return None + + +def clear_actors() -> None: + """清空注册表(测试用)。""" + _registry.clear() + _handle_cache.clear() + _background_tasks.clear() + + +def spawn_background(coro: "Awaitable[Any]") -> "asyncio.Task[Any]": + """在当前事件循环中后台调度一个协程,并维持强引用防止被 GC。 + + 等价于旧 Ray `.remote()` 的 fire-and-forget 语义: + - 协程被调度执行; + - 异常不会静默消失,而是记录到日志; + - 调用方无需 await,也无需保存返回值。 + """ + from kilostar.utils.logger import get_logger + + task: asyncio.Task[Any] = asyncio.ensure_future(coro) + _background_tasks.add(task) + + def _on_done(t: "asyncio.Task[Any]") -> None: + _background_tasks.discard(t) + if not t.cancelled() and t.exception() is not None: + get_logger("actor").exception( + f"后台任务 {t.get_name()!r} 异常退出", exc_info=t.exception() + ) + + task.add_done_callback(_on_done) + return task + + +# ─── 类型化 getter(核心单例,语义糖)────────────────────────── + + +def get_postgres() -> ActorHandle: + return get_actor("postgres_database") + + +def get_gsm() -> ActorHandle: + return get_actor("global_state_machine") + + +def get_gwm() -> ActorHandle: + return get_actor("global_workflow_manager") + + +def get_regulatory() -> ActorHandle: + return get_actor("regulatory_node") + + +def get_consciousness() -> ActorHandle: + return get_actor("consciousness_node") + + +def get_worker_cluster() -> ActorHandle: + return get_actor("worker_cluster") + + +def get_plugin_manager() -> ActorHandle: + return get_actor("global_plugin_manager") diff --git a/kilostar/utils/agent_model.py b/kilostar/utils/agent_model.py deleted file mode 100644 index 8baeae5..0000000 --- a/kilostar/utils/agent_model.py +++ /dev/null @@ -1,40 +0,0 @@ -# Copyright 2026 zhaoxi826 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - - -from pydantic import BaseModel - - -class ResponseModel(BaseModel): - """ - ResponseModel类 - 继承自pydantic的BaseModel类,是一切回复模型的父类 - """ - pass - - -class DepsModel(BaseModel): - """ - DepsModel类 - 继承自pydantic的BaseModel类,是agent运行时依赖模型的父类 - """ - pass - - -class RequestModel(BaseModel): - """ - RequestModel类 - 继承自pydantic的BaseModel类,是一切请求模型的父类 - """ - pass diff --git a/kilostar/utils/get_tool.py b/kilostar/utils/get_tool.py deleted file mode 100644 index f2fc912..0000000 --- a/kilostar/utils/get_tool.py +++ /dev/null @@ -1,117 +0,0 @@ -import importlib.util -import json -import os -import sys -from typing import Callable, Dict, List, Optional - -from kilostar.utils.logger import get_logger -from kilostar.utils.settings import get_toolset_dir - -logger = get_logger("get_tool") -_tool_cache: Dict[str, Callable] = {} -_manifest_cache: Optional[Dict[str, Dict]] = None - - -def _load_manifests() -> Dict[str, Dict]: - """扫描所有 toolset 的 manifest.json,建立 tool_name → {toolset_dir, file} 的映射。""" - global _manifest_cache - if _manifest_cache is not None: - return _manifest_cache - - _manifest_cache = {} - toolset_root = get_toolset_dir() - if not toolset_root.exists(): - return _manifest_cache - - for item in toolset_root.iterdir(): - if not item.is_dir() or item.name.startswith("__"): - continue - manifest_path = item / "manifest.json" - if not manifest_path.exists(): - continue - try: - with open(manifest_path, "r", encoding="utf-8") as f: - manifest = json.load(f) - for tool in manifest.get("tools", []): - tool_name = tool.get("name") - if tool_name: - _manifest_cache[tool_name] = { - "toolset_dir": str(item), - "toolset_name": item.name, - "file": tool.get("file", f"{tool_name}.py"), - } - except Exception as e: - logger.error(f"Failed to read manifest {manifest_path}: {e}") - - return _manifest_cache - - -def _get_tool_func(tool_name: str) -> Callable | None: - """按名字从 toolset 中加载工具函数。 - - 根据 manifest 找到工具所在的 toolset 和文件,动态加载模块并取出同名函数。 - """ - func = _tool_cache.get(tool_name) - if func: - return func - - manifests = _load_manifests() - info = manifests.get(tool_name) - if not info: - logger.error(f"Tool '{tool_name}' not found in any toolset manifest") - return None - - tool_file = os.path.join(info["toolset_dir"], info["file"]) - if not os.path.exists(tool_file): - logger.error(f"Tool file not found: {tool_file}") - return None - - try: - module_name = f"data.toolset.{info['toolset_name']}.{tool_name}" - spec = importlib.util.spec_from_file_location(module_name, tool_file) - if spec is None or spec.loader is None: - logger.error(f"Failed to create spec for {module_name}") - return None - - module = importlib.util.module_from_spec(spec) - sys.modules[module_name] = module - spec.loader.exec_module(module) - - func = getattr(module, tool_name, None) - - if not callable(func): - logger.error( - f"Tool function '{tool_name}' not found or not callable in {module_name}" - ) - return None - _tool_cache[tool_name] = func - return func - except Exception as e: - logger.error(f"Failed to load module {tool_name}: {e}") - return None - - -def del_tool_cache(tool_name: str) -> None: - """从内存缓存中移除某个工具,下次调用 ``load_tools_from_list`` 会重新从磁盘加载。""" - if tool_name in _tool_cache: - del _tool_cache[tool_name] - - -def invalidate_manifest_cache() -> None: - """清除 manifest 缓存,下次加载时重新扫描磁盘。""" - global _manifest_cache - _manifest_cache = None - - -def load_tools_from_list(tool_names: List[str] | None) -> List[Callable]: - """批量加载工具:传入工具名列表,返回成功加载到的函数对象列表(失败项被跳过)。""" - if not tool_names: - return [] - - tool_list = [] - for tool_name in tool_names: - tool_func = _get_tool_func(tool_name) - if tool_func: - tool_list.append(tool_func) - - return tool_list diff --git a/kilostar/utils/mcp_helper.py b/kilostar/utils/mcp_helper.py index fa2edea..0d06901 100644 --- a/kilostar/utils/mcp_helper.py +++ b/kilostar/utils/mcp_helper.py @@ -17,7 +17,7 @@ from typing import Dict, List, Any, Optional, Sequence from kilostar.utils.logger import get_logger -from kilostar.utils.ray_hook import ray_actor_hook +from kilostar.utils.actor import get_actor logger = get_logger("mcp_helper") @@ -132,8 +132,8 @@ async def get_all_tools_and_toolsets_for_scope( # 合入重型插件的 dispatch tools try: - pm = ray_actor_hook("global_plugin_manager").global_plugin_manager - dispatch = await pm.get_dispatch_tools.remote() + pm = get_actor("global_plugin_manager") + dispatch = await pm.get_dispatch_tools() if dispatch: tools.extend(dispatch.values()) except Exception as e: @@ -154,8 +154,8 @@ async def get_retrieval_toolsets_for_scope(scope: str) -> List[Any]: """仅返回 retrieval 工具集(system_node 专用)。不含 generation 和 MCP 工具。""" toolsets: List[Any] = [] try: - gsm = ray_actor_hook("global_state_machine").global_state_machine - retrieval = await gsm.get_retrieval_toolsets_for_scope.remote(scope) + gsm = get_actor("global_state_machine") + retrieval = await gsm.get_retrieval_toolsets_for_scope(scope) if retrieval: toolsets.extend(retrieval) except Exception as e: diff --git a/kilostar/utils/ray_compat.py b/kilostar/utils/ray_compat.py deleted file mode 100644 index f69bc14..0000000 --- a/kilostar/utils/ray_compat.py +++ /dev/null @@ -1,106 +0,0 @@ -"""KiloStar Ray 兼容层:单机/分布式模式无感切换 + 序列化工具。 - -单机模式下,所有 Actor 退化为普通 Python 异步单例,通过 StandaloneProxy -包装后暴露与 Ray Actor Handle 相同的 `.method.remote(args)` 调用接口, -使上层代码在两种模式间无感切换。 -""" - -from __future__ import annotations - -import asyncio -import os -from typing import Any, Type, TypeVar - -from pydantic import BaseModel - -_STANDALONE = os.environ.get("KILOSTAR_MODE", "distributed") == "standalone" - -T = TypeVar("T", bound=Type[BaseModel]) - - -class _MethodProxy: - """包装单个方法,使 .remote(*args, **kwargs) 返回一个可 await 的 Task。""" - - __slots__ = ("_method",) - - def __init__(self, method: Any): - self._method = method - - def remote(self, *args: Any, **kwargs: Any) -> asyncio.Task: - async def _invoke(): - result = self._method(*args, **kwargs) - if asyncio.iscoroutine(result): - return await result - return result - - return asyncio.ensure_future(_invoke()) - - -class StandaloneProxy: - """包装一个普通 Python 实例,模拟 Ray Actor Handle 的属性访问接口。 - - 用法:proxy.some_method.remote(x, y) → 等效于 await instance.some_method(x, y) - """ - - __slots__ = ("_instance",) - - def __init__(self, instance: Any): - object.__setattr__(self, "_instance", instance) - - def __getattr__(self, name: str) -> _MethodProxy: - attr = getattr(object.__getattribute__(self, "_instance"), name) - if callable(attr): - return _MethodProxy(attr) - return attr - - -# ─── 条件装饰器 ─── - - -def actor_class(cls): - """条件装饰器:分布式模式 → @ray.remote,单机模式 → 原样返回类。""" - if _STANDALONE: - return cls - import ray - return ray.remote(cls) - - -def remote_task(func): - """条件装饰器:分布式 → @ray.remote(func),单机 → .remote() 转为 asyncio task。 - - 单机模式下返回一个 stub 对象,其 .remote() 方法把函数以协程方式调度到 - 当前事件循环(workflow task 需要用 await 版本的 _entry,由调用方处理)。 - """ - if _STANDALONE: - - class _TaskProxy: - @staticmethod - def remote(*args, **kwargs): - async def _run(): - result = func(*args, **kwargs) - if asyncio.iscoroutine(result): - return await result - return result - - return asyncio.ensure_future(_run()) - - return _TaskProxy() - - import ray - return ray.remote(func) - - -# ─── Pickle (Ray 序列化优化) ─── - - -def pickle(cls: T) -> T: - """类装饰器:用 Pydantic 的高效 JSON 序列化替代 Python 原生 __reduce__, - 使 Ray 跨进程通信时对 BaseModel 子类走 Rust 级序列化。 - """ - - def __reduce__(self): - data = self.model_dump_json() - return cls.model_validate_json, (data,) - - cls.__reduce__ = __reduce__ - return cls diff --git a/kilostar/utils/ray_hook.py b/kilostar/utils/ray_hook.py deleted file mode 100644 index cc2ebed..0000000 --- a/kilostar/utils/ray_hook.py +++ /dev/null @@ -1,142 +0,0 @@ -# Copyright 2026 zhaoxi826 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -import time -from functools import lru_cache -from typing import Any, Dict - -from kilostar.utils.ray_compat import _STANDALONE - -if not _STANDALONE: - import ray - - -class ActorList: - """属性式访问的简易容器,用 ``a.actor_name`` 取代 ``d["actor_name"]``。""" - - def __init__(self): - super().__setattr__("dict", {}) - - def __setattr__(self, key, value): - self.dict[key] = value - - def __getattr__(self, key): - if key in self.dict: - return self.dict[key] - raise AttributeError(f"ActorList 对象没有属性 '{key}'") - - def __delattr__(self, key): - if key in self.dict: - del self.dict[key] - else: - raise AttributeError(f"ActorList对象没有属性 '{key}'") - - -# ─── Standalone Registry ─── - -_standalone_registry: Dict[str, Any] = {} - - -def register_standalone(name: str, instance: Any) -> None: - """注册一个单机模式下的 Actor 单例(已包装为 StandaloneProxy)。""" - from kilostar.utils.ray_compat import StandaloneProxy - - _standalone_registry[name] = StandaloneProxy(instance) - - -# ─── Distributed Mode Helpers ─── - - -if not _STANDALONE: - - @lru_cache(maxsize=128) - def _get_cached_actor_handle(actor_name: str): - """缓存接口""" - return ray.get_actor(actor_name, namespace="kilostar") - - def clear_actor_cache(): - """清理接口""" - _get_cached_actor_handle.cache_clear() - - def wait_for_actor( - actor_name: str, *, timeout: float = 10.0, interval: float = 0.5 - ): - """阻塞等待某个 actor 就绪,返回其句柄。""" - deadline = time.monotonic() + max(timeout, 0.0) - last_err: Exception | None = None - while True: - try: - return _get_cached_actor_handle(actor_name) - except Exception as e: - last_err = e - if time.monotonic() >= deadline: - raise TimeoutError( - f"等待 actor {actor_name!r} 就绪超时({timeout}s):{last_err}" - ) from last_err - time.sleep(interval) - -else: - - def _get_cached_actor_handle(actor_name: str): - raise RuntimeError("单机模式下不应调用 _get_cached_actor_handle") - - def clear_actor_cache(): - pass - - def wait_for_actor(actor_name: str, **kwargs): - raise RuntimeError("单机模式下不应调用 wait_for_actor") - - -# ─── 统一入口 ─── - - -def ray_actor_hook(*actor_names: str, timeout: float = 0.0, interval: float = 0.5): - """按名字批量取出 Actor 句柄,组装成一个 ActorList 返回。 - - 单机模式从 _standalone_registry 取,分布式模式走 ray.get_actor。 - """ - actor_list = ActorList() - - if _STANDALONE: - for name in actor_names: - if name not in _standalone_registry: - raise ValueError( - f"Standalone registry: actor {name!r} not registered" - ) - setattr(actor_list, name, _standalone_registry[name]) - return actor_list - - for actor_name in actor_names: - if timeout > 0: - handle = wait_for_actor( - actor_name, timeout=timeout, interval=interval - ) - else: - handle = _get_cached_actor_handle(actor_name) - setattr(actor_list, actor_name, handle) - return actor_list - - -def get_worker_cluster(affinity: str = "cpu"): - """按 node_affinity 标签取对应的 WorkerCluster actor 句柄。 - - 单机模式统一返回唯一的 worker_cluster 实例。 - 分布式模式按 affinity 路由到 worker_cluster_cpu / _core / _gpu。 - 未知标签降级到 cpu。 - """ - if _STANDALONE: - return _standalone_registry.get("worker_cluster") - - _valid = {"cpu", "core", "gpu"} - node_type = affinity if affinity in _valid else "cpu" - return _get_cached_actor_handle(f"worker_cluster_{node_type}") diff --git a/kilostar/utils/request_context.py b/kilostar/utils/request_context.py index d0592ec..03e96a1 100644 --- a/kilostar/utils/request_context.py +++ b/kilostar/utils/request_context.py @@ -27,8 +27,8 @@ 1. ``contextvars`` 在 ``asyncio`` 协程间天然继承,不会跨协程串味; 2. ``loguru`` 的 ``patcher`` 钩子可以把它变成日志切面,业务代码不需要在每条 ``logger.info`` 上手动 ``.bind(trace_id=...)``; -3. Ray 跨进程调用时 contextvars 不会自动传播 —— 这是有意为之,避免不同 actor - 间的上下文意外串联。跨 actor 边界要走显式参数,由接收方再 ``bind_*`` 一次。 +3. 跨 asyncio.create_task 边界时也会自动继承当前 context,符合单进程 actor 模型 + 的语义 —— 一次任务链路的 ID 与其入口请求一致,不必显式透传。 """ from __future__ import annotations @@ -36,7 +36,7 @@ from __future__ import annotations import uuid from contextlib import contextmanager from contextvars import ContextVar, Token -from typing import Iterator, Optional +from typing import Iterator _request_id_var: ContextVar[str] = ContextVar("kilostar_request_id", default="") @@ -101,30 +101,8 @@ def new_request_id(prefix: str = "req") -> str: def snapshot() -> dict[str, str]: - """返回当前上下文 ID 的快照,便于跨 actor/task 边界显式透传。""" + """返回当前上下文 ID 的快照,便于日志/追踪或结构化写库。""" return { "request_id": _request_id_var.get(), "trace_id": _trace_id_var.get(), } - - -@contextmanager -def apply_snapshot(snap: Optional[dict[str, str]]) -> Iterator[None]: - """把外部传来的 snapshot 在当前 context 内生效一次(用于跨 Ray actor 调用时)。""" - if not snap: - yield - return - tokens: list[Token] = [] - if snap.get("request_id"): - tokens.append(_request_id_var.set(snap["request_id"])) - if snap.get("trace_id"): - tokens.append(_trace_id_var.set(snap["trace_id"])) - try: - yield - finally: - for tok in reversed(tokens): - try: - tok.var.reset(tok) - except (ValueError, LookupError): - # token 可能因协程切换失效,宽容处理 - pass diff --git a/kilostar/utils/settings.py b/kilostar/utils/settings.py index 6ef3b65..ea4fb97 100644 --- a/kilostar/utils/settings.py +++ b/kilostar/utils/settings.py @@ -38,7 +38,6 @@ class OnebotSettings(BaseSettings): class AppSettings(BaseSettings): - kilostar_mode: str = "distributed" kilostar_lang: str = "zh" kilostar_cors_origins: str = "" kilostar_plugin_dir: str = "" diff --git a/kilostar/worker_cluster/worker_cluster.py b/kilostar/worker_cluster/worker_cluster.py index 188b771..e1f17e7 100644 --- a/kilostar/worker_cluster/worker_cluster.py +++ b/kilostar/worker_cluster/worker_cluster.py @@ -15,13 +15,7 @@ import time import asyncio from collections import OrderedDict -from kilostar.utils.ray_compat import actor_class, _STANDALONE -from kilostar.utils.ray_hook import ray_actor_hook - -if _STANDALONE: - from asyncio import Queue -else: - from ray.util.queue import Queue +from asyncio import Queue from kilostar.worker_individual.base_individual import BaseIndividual from kilostar.worker_individual.skill_individual import SkillIndividual from kilostar.worker_individual.ordinary_individual import OrdinaryIndividual @@ -31,15 +25,14 @@ from kilostar.worker_individual.special_individual import SpecialIndividual from kilostar.utils.logger import get_logger -@actor_class class WorkerCluster: """ - 工作集群 Actor:管理和调度所有的 worker_individual - 设计理念:按需加载,内存 LRU 淘汰,避免 Actor 爆炸 + 工作集群:管理和调度所有的 worker_individual + 设计理念:按需加载,内存 LRU 淘汰,避免实例爆炸 - 分布式模式下每种 node_type 对应一个独立实例,Ray 根据自定义资源 - ``kilostar_node_cpu`` / ``kilostar_node_core`` / ``kilostar_node_gpu`` - 将 Actor 调度到声明了对应资源的节点上。 + ``submit_task`` 是对外的稳定执行边界——未来若要把重型 worker(如 vLLM 本地 + 推理)拆到独立进程/机器上,只需在这个边界后面换一个 executor 实现,上层不感知。 + ``node_type`` 字段暂时保留,供未来按算力亲和性路由使用。 """ def __init__(self, max_capacity: int = 200, num_runners: int = 10, node_type: str = "cpu"): @@ -70,11 +63,8 @@ class WorkerCluster: from kilostar.core.global_state_machine.gsm_snapshot import fetch_snapshot - global_state_machine = ray_actor_hook( - "global_state_machine" - ).global_state_machine - # 走快照读,避开 GSM actor RPC:高频唤醒路径不再是单 actor 瓶颈 - snapshot = await fetch_snapshot(gsm_actor=global_state_machine) + # 走快照读:高频唤醒路径直接读进程内缓存快照 + snapshot = await fetch_snapshot() agent_config = snapshot.individuals.get(agent_id) if not agent_config: @@ -104,7 +94,7 @@ class WorkerCluster: if self.task_queue is None: await asyncio.sleep(0.1) continue - task = await self.task_queue.get() if _STANDALONE else await self.task_queue.get_async() + task = await self.task_queue.get() task_id = task.get("task_id") agent_id = task.get("agent_id") task_event = task.get("task_event") @@ -150,10 +140,7 @@ class WorkerCluster: self.results_futures[task_id] = future task = {"task_id": task_id, "agent_id": agent_id, "task_event": task_event} - if _STANDALONE: - await self.task_queue.put(task) - else: - await self.task_queue.put_async(task) + await self.task_queue.put(task) self.logger.debug(f"[WorkerCluster] 任务 {task_id} 已加入队列。") try: @@ -168,5 +155,5 @@ class WorkerCluster: "active_worker_count": len(self._active_workers), "max_capacity": self.max_capacity, "cached_agent_ids": list(self._active_workers.keys()), - "queue_size": self.task_queue.qsize() if _STANDALONE else self.task_queue.size(), + "queue_size": self.task_queue.qsize() if self.task_queue else 0, } diff --git a/kilostar/worker_individual/base_individual.py b/kilostar/worker_individual/base_individual.py index c2997d0..ad18c94 100644 --- a/kilostar/worker_individual/base_individual.py +++ b/kilostar/worker_individual/base_individual.py @@ -13,30 +13,28 @@ # limitations under the License. from pydantic_ai import Agent, RunContext -from pydantic import Field +from pydantic import BaseModel, Field from kilostar.adapter.model_adapter.agent_factory import AgentFactory from kilostar.core.global_state_machine.model_provider.base_provider import Provider -from kilostar.utils.agent_model import ResponseModel, RequestModel, DepsModel -from kilostar.utils.ray_hook import ray_actor_hook from kilostar.utils.logger import get_logger logger = get_logger("worker_individual") -class WorkerIndividualResponse(ResponseModel): +class WorkerIndividualResponse(BaseModel): """Worker Individual 的输出模型,承载一次任务执行后的结果文本。""" output: str = Field(..., description="Worker执行任务的输出结果") -class WorkerIndividualDeps(DepsModel): +class WorkerIndividualDeps(BaseModel): """Worker Individual 的运行期依赖,注入到 pydantic-ai Agent 的 RunContext。""" task_event: dict -class WorkerIndividualInput(RequestModel): +class WorkerIndividualInput(BaseModel): """Worker Individual 的输入模型,承载一次任务事件的入参。""" task_event: dict @@ -67,17 +65,14 @@ class BaseIndividual: from kilostar.utils.mcp_helper import get_all_tools_and_toolsets_for_scope from kilostar.core.global_state_machine.gsm_snapshot import fetch_snapshot - global_state_machine = ray_actor_hook( - "global_state_machine" - ).global_state_machine provider_title = self.agent_config.get( "provider_title", "openai" ) # default fallback model_id = self.agent_config.get("model_id", "gpt-4o") # default fallback toolset_ids = self.agent_config.get("tools", None) - # 直读快照,避开 actor RPC 单线程串行 - snapshot = await fetch_snapshot(gsm_actor=global_state_machine) + # 直读进程内缓存快照 + snapshot = await fetch_snapshot() provider: Provider = snapshot.providers.get(provider_title) if provider is None: raise ValueError(f"Provider {provider_title!r} 未注册") diff --git a/main.py b/main.py index 272228e..6620fe1 100644 --- a/main.py +++ b/main.py @@ -32,8 +32,6 @@ except Exception as e: import asyncio -KILOSTAR_MODE = os.environ.get("KILOSTAR_MODE", "distributed") - from kilostar.worker_cluster import WorkerCluster from kilostar.utils.banner import print_banner from kilostar.core.postgres_database import PostgresDatabase @@ -43,155 +41,52 @@ from kilostar.core.individual.regulatory_node import RegulatoryNode from kilostar.core.individual.consciousness_node import ConsciousnessNode from kilostar.plugin_runtime.plugin_manager import GlobalPluginManager -if KILOSTAR_MODE != "standalone": - import ray - from ray import serve - from kilostar.api import KiloStarGateway +async def start() -> None: + """启动 KiloStar:单进程 asyncio 单体。 -async def start_standalone(): - """单机模式:纯 asyncio,不依赖 Ray。""" + 所有核心组件在本进程内实例化,经 ``register_actor`` 登记到进程内 actor 运行时; + 业务代码统一通过 ``get_actor(name)`` 拿句柄调用(位置透明,未来可换远程后端)。 + """ import uvicorn - from kilostar.utils.ray_hook import register_standalone + from kilostar.utils.actor import register_actor, get_actor from kilostar.api import app postgres_database = PostgresDatabase() await postgres_database.init_db() - register_standalone("postgres_database", postgres_database) + register_actor("postgres_database", postgres_database) - from kilostar.utils.ray_compat import StandaloneProxy - postgres_proxy = StandaloneProxy(postgres_database) - - global_state_machine = GlobalStateMachine(postgres_proxy) + # GSM 通过 actor 句柄依赖 postgres:保持"依赖对端不关心其位置"的一致语义 + global_state_machine = GlobalStateMachine(get_actor("postgres_database")) await global_state_machine.init_state_machine() - register_standalone("global_state_machine", global_state_machine) + register_actor("global_state_machine", global_state_machine) global_workflow_manager = GlobalWorkflowManager() await global_workflow_manager.init_manager() - register_standalone("global_workflow_manager", global_workflow_manager) + register_actor("global_workflow_manager", global_workflow_manager) - regulatory_node = RegulatoryNode() - register_standalone("regulatory_node", regulatory_node) - - consciousness_node = ConsciousnessNode() - register_standalone("consciousness_node", consciousness_node) + register_actor("regulatory_node", RegulatoryNode()) + register_actor("consciousness_node", ConsciousnessNode()) worker_cluster = WorkerCluster(node_type="cpu") await worker_cluster.start() - register_standalone("worker_cluster", worker_cluster) - # 单机模式三个标签共用同一实例 - register_standalone("worker_cluster_cpu", worker_cluster) - register_standalone("worker_cluster_core", worker_cluster) - register_standalone("worker_cluster_gpu", worker_cluster) + register_actor("worker_cluster", worker_cluster) plugin_manager = GlobalPluginManager() await plugin_manager.bootstrap() - register_standalone("global_plugin_manager", plugin_manager) + register_actor("global_plugin_manager", plugin_manager) - print(f"✅ KiloStar 单机模式启动完成,监听 0.0.0.0:8000") + print("✅ KiloStar 启动完成,监听 0.0.0.0:8000") config = uvicorn.Config(app, host="0.0.0.0", port=8000, log_level="info") server = uvicorn.Server(config) await server.serve() -async def start_distributed(): - """分布式模式:使用 Ray Actor + Ray Serve。""" - env_vars = { - "POSTGRES_USER": os.getenv("POSTGRES_USER", "postgres"), - "POSTGRES_PASSWORD": os.getenv("POSTGRES_PASSWORD", ""), - "POSTGRES_HOST": os.getenv("POSTGRES_HOST", "db"), - "POSTGRES_PORT": os.getenv("POSTGRES_PORT", "5432"), - "POSTGRES_DB": os.getenv("POSTGRES_DB", "postgres"), - "SECRET_KEY": os.getenv("SECRET_KEY"), - } - - ray.init( - ignore_reinit_error=True, - namespace="kilostar", - dashboard_host="0.0.0.0", - dashboard_port=8265, - runtime_env={"env_vars": env_vars}, - resources={ - "kilostar_node_cpu": 1, - "kilostar_node_core": 1, - "kilostar_node_gpu": 1, - }, - ) - - postgres_database = PostgresDatabase.options( - name="postgres_database" - ).remote() - await postgres_database.init_db.remote() - - global_state_machine = GlobalStateMachine.options( - name="global_state_machine", namespace="kilostar", lifetime="detached" - ).remote(postgres_database) - - print("正在等待 GlobalStateMachine 初始化并加载注册表...") - try: - await global_state_machine.init_state_machine.remote() - print("GlobalStateMachine 初始化成功!") - except Exception as e: - print(f"\n[致命错误] GlobalStateMachine 启动失败!\n{e}\n") - return - - global_workflow_manager = GlobalWorkflowManager.options( - name="global_workflow_manager", namespace="kilostar", lifetime="detached" - ).remote() - - RegulatoryNode.options(name="regulatory_node").remote() - ConsciousnessNode.options(name="consciousness_node").remote() - - try: - for node_type in ("cpu", "core", "gpu"): - actor_name = f"worker_cluster_{node_type}" - resource_key = f"kilostar_node_{node_type}" - try: - WorkerCluster.options( - name=actor_name, - lifetime="detached", - resources={resource_key: 1}, - ).remote(node_type=node_type) - print(f"✅ WorkerCluster[{node_type}] 已成功启动并注册!") - except ValueError: - print(f"WorkerCluster[{node_type}] 已经存在。") - except Exception as e: - print(f"WorkerCluster 启动失败: {e}") - - print("正在等待 GlobalWorkflowManager 初始化与恢复工作流...") - try: - await global_workflow_manager.init_manager.remote() - print("GlobalWorkflowManager 初始化成功!") - except Exception as e: - print(f"\n[致命错误] GlobalWorkflowManager 启动失败!\n{e}\n") - return - - plugin_manager = GlobalPluginManager.options( - name="global_plugin_manager", namespace="kilostar", lifetime="detached" - ).remote() - try: - await plugin_manager.bootstrap.remote() - print("✅ GlobalPluginManager 初始化成功!") - except Exception as e: - print(f"⚠️ GlobalPluginManager 启动失败(非致命): {e}") - - serve.start(http_options={"host": "0.0.0.0", "port": 8000}) - serve.run(KiloStarGateway.bind()) - - while True: - await asyncio.sleep(3600) - - def main(): print_banner() - mode = KILOSTAR_MODE - print(f"启动模式: {mode}") try: - if mode == "standalone": - asyncio.run(start_standalone()) - else: - asyncio.run(start_distributed()) + asyncio.run(start()) except KeyboardInterrupt: print("系统已退出。") diff --git a/pyproject.toml b/pyproject.toml index 9b8f261..f46ef04 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -17,6 +17,7 @@ dependencies = [ "alembic>=1.13.0", "asyncpg>=0.31.0", "cryptography>=42.0.0", + "fastapi>=0.115.0", "httpx>=0.28.1", "jinja2>=3.1.6", "loguru>=0.7.3", @@ -27,10 +28,10 @@ dependencies = [ "pyfiglet>=1.0.4", "pyjwt>=2.12.1", "python-ulid>=3.1.0", - "ray[default,serve]>=2.54.0", "rich>=14.3.3", "sqlalchemy>=2.0.49", "tavily-python>=0.7.0", + "uvicorn[standard]>=0.30.0", ] [project.optional-dependencies] diff --git a/tests/conftest.py b/tests/conftest.py index ab575cd..cddb3f4 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,76 +1,56 @@ -"""Pytest 全局 fixture:把 Ray Actor 句柄、PostgresDatabase、loguru 等重副作用模块替换成可控的 stub。""" +"""Pytest 全局 fixture:把 actor 句柄、PostgresDatabase、loguru 等重副作用替换成可控 stub。""" from __future__ import annotations -import sys import types -from typing import Any, Dict, Optional +from typing import Any, Dict from unittest.mock import AsyncMock, MagicMock import pytest # ───────────────────────────────────────────────────────────────────────────── -# Ray actor 句柄存根:测试期不真正连 Ray,由 conftest 注入名字 -> AsyncMock 句柄 +# actor 注册表存根:测试期把名字 -> mock 实例塞进进程内 actor 运行时。 +# 生产代码 ``get_actor(name).method(args)`` 会拿到 LocalActorHandle 转发到 mock。 +# mock 方法用 AsyncMock 即可(handle 统一返回 awaitable)。 # ───────────────────────────────────────────────────────────────────────────── class _FakeActorRegistry: - """模拟 ``ray.get_actor`` 行为:测试可往里塞名字 -> AsyncMock。""" + """薄封装:``register(name, instance)`` 直接登记到真实 actor 运行时。""" - def __init__(self) -> None: - self._actors: Dict[str, Any] = {} + def register(self, name: str, instance: Any) -> None: + from kilostar.utils import actor - def register(self, name: str, handle: Any) -> None: - self._actors[name] = handle - - def get(self, name: str, namespace: str = "kilostar"): # noqa: ARG002 - if name not in self._actors: - raise ValueError(f"FakeActorRegistry: actor {name!r} not registered") - return self._actors[name] - - def clear(self) -> None: - self._actors.clear() + actor.register_actor(name, instance) @pytest.fixture def fake_actors(monkeypatch) -> _FakeActorRegistry: - """把 ``kilostar.utils.ray_hook._get_cached_actor_handle`` 的实现替换为 fake registry。 + """清空并接管进程内 actor 运行时,测试结束后复位。 用法:: def test_xxx(fake_actors): - gsm = AsyncMock() - gsm.get_tool_config.remote = AsyncMock(return_value={"api_key": "k"}) - fake_actors.register("global_state_machine", types.SimpleNamespace(global_state_machine=gsm)) + gsm = MagicMock() + gsm.get_tool_config = AsyncMock(return_value={"api_key": "k"}) + fake_actors.register("global_state_machine", gsm) """ - registry = _FakeActorRegistry() + from kilostar.utils import actor - from kilostar.utils import ray_hook - - ray_hook.clear_actor_cache() - original = ray_hook._get_cached_actor_handle - - def _stub(actor_name: str): - return registry.get(actor_name) - - _stub.cache_clear = lambda: None # type: ignore[attr-defined] - monkeypatch.setattr(ray_hook, "_get_cached_actor_handle", _stub) - yield registry - registry.clear() - monkeypatch.setattr(ray_hook, "_get_cached_actor_handle", original) - ray_hook.clear_actor_cache() + actor.clear_actors() + yield _FakeActorRegistry() + actor.clear_actors() @pytest.fixture def gsm_handle(fake_actors) -> MagicMock: - """快捷 fixture:注册一个名为 ``global_state_machine`` 的 actor,返回其内部 mock。 + """快捷 fixture:注册名为 ``global_state_machine`` 的 mock 实例并返回它。 - 内部 mock 的方法默认全部是 ``AsyncMock``,调用 ``.remote(...)`` 会按 AsyncMock 规则返回。 + 方法默认为 ``AsyncMock``;生产侧 ``await get_actor(...).method(args)`` 正常工作。 """ gsm = MagicMock() - container = types.SimpleNamespace(global_state_machine=gsm) - fake_actors.register("global_state_machine", container) + fake_actors.register("global_state_machine", gsm) return gsm diff --git a/tests/unit/test_agent_factory.py b/tests/unit/test_agent_factory.py index 34a9788..8bad668 100644 --- a/tests/unit/test_agent_factory.py +++ b/tests/unit/test_agent_factory.py @@ -14,9 +14,13 @@ import pytest from kilostar.adapter.model_adapter import agent_factory as af_mod from kilostar.adapter.model_adapter.agent_factory import AgentFactory from kilostar.core.global_state_machine.model_provider.base_provider import Provider -from kilostar.utils.agent_model import DepsModel, ResponseModel +from pydantic import BaseModel as _BaseModel + from kilostar.utils.error import ModelNotExistError +ResponseModel = _BaseModel +DepsModel = _BaseModel + class _SpyProvider: last_init: Dict[str, Any] = {} diff --git a/tests/unit/test_api_agent_template.py b/tests/unit/test_api_agent_template.py index 50c2289..126f5e9 100644 --- a/tests/unit/test_api_agent_template.py +++ b/tests/unit/test_api_agent_template.py @@ -85,7 +85,7 @@ def test_update_none_affinity_ok(): @pytest.mark.asyncio async def test_list_templates(app, fake_actors): pg = types.SimpleNamespace( - list_templates=types.SimpleNamespace(remote=AsyncMock(return_value=[])) + list_templates=AsyncMock(return_value=[]) ) fake_actors.register("postgres_database", pg) async with AsyncClient(transport=ASGITransport(app=app), base_url="http://t") as c: @@ -98,7 +98,7 @@ async def test_list_templates(app, fake_actors): async def test_create_template(app, fake_actors): tpl = _tpl() pg = types.SimpleNamespace( - add_template=types.SimpleNamespace(remote=AsyncMock(return_value=tpl)) + add_template=AsyncMock(return_value=tpl) ) fake_actors.register("postgres_database", pg) async with AsyncClient(transport=ASGITransport(app=app), base_url="http://t") as c: @@ -110,7 +110,7 @@ async def test_create_template(app, fake_actors): @pytest.mark.asyncio async def test_delete_template_not_found(app, fake_actors): pg = types.SimpleNamespace( - get_template=types.SimpleNamespace(remote=AsyncMock(return_value=None)) + get_template=AsyncMock(return_value=None) ) fake_actors.register("postgres_database", pg) async with AsyncClient(transport=ASGITransport(app=app), base_url="http://t") as c: @@ -121,7 +121,7 @@ async def test_delete_template_not_found(app, fake_actors): @pytest.mark.asyncio async def test_delete_other_users_template_forbidden(app, fake_actors): pg = types.SimpleNamespace( - get_template=types.SimpleNamespace(remote=AsyncMock(return_value=_tpl(owner="bob"))) + get_template=AsyncMock(return_value=_tpl(owner="bob")) ) fake_actors.register("postgres_database", pg) async with AsyncClient(transport=ASGITransport(app=app), base_url="http://t") as c: diff --git a/tests/unit/test_api_chat.py b/tests/unit/test_api_chat.py index bde5cf0..c4465da 100644 --- a/tests/unit/test_api_chat.py +++ b/tests/unit/test_api_chat.py @@ -7,22 +7,13 @@ from __future__ import annotations -from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock import pytest -class _FakeActorRef: - def __init__(self, target): - self._target = target - - def __getattr__(self, item): - return getattr(self._target, item) - - @pytest.mark.asyncio -async def test_ask_regulatory_calls_working_and_extracts_reply(monkeypatch): +async def test_ask_regulatory_calls_working_and_extracts_reply(fake_actors): from kilostar.api import chat as chat_module from kilostar.core.individual.regulatory_node.template import MessageResponse @@ -31,14 +22,8 @@ async def test_ask_regulatory_calls_working_and_extracts_reply(monkeypatch): ) regulatory = MagicMock() - regulatory.working = MagicMock() - regulatory.working.remote = AsyncMock(return_value=fake_resp) - - def _fake_hook(name): - assert name == "regulatory_node" - return SimpleNamespace(regulatory_node=_FakeActorRef(regulatory)) - - monkeypatch.setattr(chat_module, "ray_actor_hook", _fake_hook) + regulatory.working = AsyncMock(return_value=fake_resp) + fake_actors.register("regulatory_node", regulatory) out = await chat_module._ask_regulatory( user_id="alice", chat_id="chat-1", message="hi" @@ -46,7 +31,7 @@ async def test_ask_regulatory_calls_working_and_extracts_reply(monkeypatch): assert out == "你好" # 调用契约:MessageRequest,且 platform_id 取自 chat_id - args, kwargs = regulatory.working.remote.call_args + args, kwargs = regulatory.working.call_args payload = args[0] if args else kwargs.get("payload") assert payload.platform == "client" assert payload.user_name == "alice" @@ -55,19 +40,13 @@ async def test_ask_regulatory_calls_working_and_extracts_reply(monkeypatch): @pytest.mark.asyncio -async def test_ask_regulatory_returns_none_when_node_returns_none(monkeypatch): +async def test_ask_regulatory_returns_none_when_node_returns_none(fake_actors): """节点降级返回 None 时,上层应静默不写回 chat history。""" from kilostar.api import chat as chat_module regulatory = MagicMock() - regulatory.working = MagicMock() - regulatory.working.remote = AsyncMock(return_value=None) - - monkeypatch.setattr( - chat_module, - "ray_actor_hook", - lambda name: SimpleNamespace(regulatory_node=_FakeActorRef(regulatory)), - ) + regulatory.working = AsyncMock(return_value=None) + fake_actors.register("regulatory_node", regulatory) out = await chat_module._ask_regulatory( user_id="bob", chat_id="chat-2", message="hello" diff --git a/tests/unit/test_api_custom_toolset_auth.py b/tests/unit/test_api_custom_toolset_auth.py index b106739..303b329 100644 --- a/tests/unit/test_api_custom_toolset_auth.py +++ b/tests/unit/test_api_custom_toolset_auth.py @@ -39,10 +39,8 @@ async def test_get_custom_toolset_forbidden_for_non_owner( app_with_user, fake_actors ): gsm = types.SimpleNamespace() - gsm.get_custom_toolset = types.SimpleNamespace( - remote=AsyncMock( - return_value={"toolset_id": "t1", "owner_id": "bob", "tools": []} - ) + gsm.get_custom_toolset = AsyncMock( + return_value={"toolset_id": "t1", "owner_id": "bob", "tools": []} ) fake_actors.register("global_state_machine", gsm) @@ -56,10 +54,8 @@ async def test_get_custom_toolset_forbidden_for_non_owner( @pytest.mark.asyncio async def test_get_custom_toolset_allowed_for_owner(app_with_user, fake_actors): gsm = types.SimpleNamespace() - gsm.get_custom_toolset = types.SimpleNamespace( - remote=AsyncMock( - return_value={"toolset_id": "t1", "owner_id": "alice", "tools": []} - ) + gsm.get_custom_toolset = AsyncMock( + return_value={"toolset_id": "t1", "owner_id": "alice", "tools": []} ) fake_actors.register("global_state_machine", gsm) @@ -76,10 +72,8 @@ async def test_get_custom_toolset_allowed_for_admin( app_with_user, fake_actors, monkeypatch ): gsm = types.SimpleNamespace() - gsm.get_custom_toolset = types.SimpleNamespace( - remote=AsyncMock( - return_value={"toolset_id": "t1", "owner_id": "bob", "tools": []} - ) + gsm.get_custom_toolset = AsyncMock( + return_value={"toolset_id": "t1", "owner_id": "bob", "tools": []} ) fake_actors.register("global_state_machine", gsm) @@ -103,9 +97,7 @@ async def test_list_custom_toolsets_filters_by_owner(app_with_user, fake_actors) {"toolset_id": "t2", "owner_id": "bob"}, ] gsm = types.SimpleNamespace() - gsm.list_custom_toolsets = types.SimpleNamespace( - remote=AsyncMock(return_value=all_sets) - ) + gsm.list_custom_toolsets = AsyncMock(return_value=all_sets) fake_actors.register("global_state_machine", gsm) transport = ASGITransport(app=app_with_user) @@ -123,13 +115,11 @@ async def test_delete_custom_toolset_forbidden_for_non_owner( app_with_user, fake_actors ): gsm = types.SimpleNamespace() - gsm.get_custom_toolset = types.SimpleNamespace( - remote=AsyncMock( - return_value={"toolset_id": "t1", "owner_id": "bob"} - ) + gsm.get_custom_toolset = AsyncMock( + return_value={"toolset_id": "t1", "owner_id": "bob"} ) delete_mock = AsyncMock(return_value=True) - gsm.delete_custom_toolset = types.SimpleNamespace(remote=delete_mock) + gsm.delete_custom_toolset = delete_mock fake_actors.register("global_state_machine", gsm) transport = ASGITransport(app=app_with_user) diff --git a/tests/unit/test_api_health.py b/tests/unit/test_api_health.py index 6c8f521..e1a8242 100644 --- a/tests/unit/test_api_health.py +++ b/tests/unit/test_api_health.py @@ -31,11 +31,11 @@ async def test_liveness_returns_alive(health_app): @pytest.mark.asyncio async def test_readiness_all_ok(health_app, fake_actors): pg = types.SimpleNamespace() - pg.ping = types.SimpleNamespace(remote=AsyncMock(return_value=True)) + pg.ping = AsyncMock(return_value=True) fake_actors.register("postgres_database", pg) gsm = types.SimpleNamespace() - gsm.get_skill_list = types.SimpleNamespace(remote=AsyncMock(return_value=[])) + gsm.get_skill_list = AsyncMock(return_value=[]) fake_actors.register("global_state_machine", gsm) transport = ASGITransport(app=health_app) @@ -51,13 +51,11 @@ async def test_readiness_all_ok(health_app, fake_actors): @pytest.mark.asyncio async def test_readiness_postgres_down(health_app, fake_actors): pg = types.SimpleNamespace() - pg.ping = types.SimpleNamespace( - remote=AsyncMock(side_effect=RuntimeError("db down")) - ) + pg.ping = AsyncMock(side_effect=RuntimeError("db down")) fake_actors.register("postgres_database", pg) gsm = types.SimpleNamespace() - gsm.get_skill_list = types.SimpleNamespace(remote=AsyncMock(return_value=[])) + gsm.get_skill_list = AsyncMock(return_value=[]) fake_actors.register("global_state_machine", gsm) transport = ASGITransport(app=health_app) diff --git a/tests/unit/test_api_onebot.py b/tests/unit/test_api_onebot.py index 8dad600..b994aa8 100644 --- a/tests/unit/test_api_onebot.py +++ b/tests/unit/test_api_onebot.py @@ -78,7 +78,7 @@ def test_extract_plain_text_handles_unknown_type(): def regulatory_actor(fake_actors): """注入一个 regulatory_node Mock,默认返回固定的 MessageResponse。""" inner = MagicMock() - inner.working.remote = AsyncMock( + inner.working = AsyncMock( return_value=MessageResponse( platform="onebot", platform_id="private:1234", reply_message="pong" ) @@ -92,7 +92,7 @@ async def test_dispatch_event_ignores_non_message(regulatory_actor): {"post_type": "meta_event", "meta_event_type": "heartbeat"} ) assert res is None - regulatory_actor.working.remote.assert_not_called() + regulatory_actor.working.assert_not_called() async def test_dispatch_event_ignores_empty_text(regulatory_actor): @@ -105,7 +105,7 @@ async def test_dispatch_event_ignores_empty_text(regulatory_actor): } ) assert res is None - regulatory_actor.working.remote.assert_not_called() + regulatory_actor.working.assert_not_called() async def test_dispatch_event_private_message_returns_quick_reply(regulatory_actor): @@ -140,7 +140,7 @@ async def test_dispatch_event_group_message_includes_at_sender(regulatory_actor) async def test_dispatch_event_swallows_actor_error(regulatory_actor): - regulatory_actor.working.remote = AsyncMock(side_effect=RuntimeError("ray fail")) + regulatory_actor.working = AsyncMock(side_effect=RuntimeError("ray fail")) payload = { "post_type": "message", "message_type": "private", @@ -153,7 +153,7 @@ async def test_dispatch_event_swallows_actor_error(regulatory_actor): async def test_dispatch_event_returns_none_when_reply_empty(fake_actors): inner = MagicMock() - inner.working.remote = AsyncMock( + inner.working = AsyncMock( return_value=MessageResponse( platform="onebot", platform_id="private:1", reply_message="" ) diff --git a/tests/unit/test_api_workflow_auth.py b/tests/unit/test_api_workflow_auth.py index 3aaf039..e7682dc 100644 --- a/tests/unit/test_api_workflow_auth.py +++ b/tests/unit/test_api_workflow_auth.py @@ -36,9 +36,9 @@ def app_alice(): def _register_pg(fake_actors, owner: str = "alice"): pg = types.SimpleNamespace() - pg.get_workflow = types.SimpleNamespace(remote=AsyncMock(return_value=_make_workflow(owner))) - pg.get_workflow_context = types.SimpleNamespace(remote=AsyncMock(return_value=None)) - pg.get_workflow_graph_state = types.SimpleNamespace(remote=AsyncMock(return_value=None)) + pg.get_workflow = AsyncMock(return_value=_make_workflow(owner)) + pg.get_workflow_context = AsyncMock(return_value=None) + pg.get_workflow_graph_state = AsyncMock(return_value=None) fake_actors.register("postgres_database", pg) return pg @@ -54,7 +54,7 @@ async def test_detail_forbidden_other_user(app_alice, fake_actors): @pytest.mark.asyncio async def test_detail_not_found(app_alice, fake_actors): pg = types.SimpleNamespace() - pg.get_workflow = types.SimpleNamespace(remote=AsyncMock(return_value=None)) + pg.get_workflow = AsyncMock(return_value=None) fake_actors.register("postgres_database", pg) async with AsyncClient(transport=ASGITransport(app=app_alice), base_url="http://t") as c: resp = await c.get("/api/v1/workflow/trace-nonexist") @@ -80,7 +80,7 @@ async def test_resume_forbidden_other_user(app_alice, fake_actors): @pytest.mark.asyncio async def test_resume_not_found(app_alice, fake_actors): pg = types.SimpleNamespace() - pg.get_workflow = types.SimpleNamespace(remote=AsyncMock(return_value=None)) + pg.get_workflow = AsyncMock(return_value=None) fake_actors.register("postgres_database", pg) async with AsyncClient(transport=ASGITransport(app=app_alice), base_url="http://t") as c: resp = await c.post("/api/v1/workflow/trace-nonexist/resume") @@ -106,7 +106,7 @@ async def test_sse_forbidden_other_user(app_alice, fake_actors): @pytest.mark.asyncio async def test_sse_not_found(app_alice, fake_actors): pg = types.SimpleNamespace() - pg.get_workflow = types.SimpleNamespace(remote=AsyncMock(return_value=None)) + pg.get_workflow = AsyncMock(return_value=None) fake_actors.register("postgres_database", pg) async with AsyncClient(transport=ASGITransport(app=app_alice), base_url="http://t") as c: resp = await c.get("/api/v1/workflow/sse/trace-nonexist") diff --git a/tests/unit/test_gsm_registries.py b/tests/unit/test_gsm_registries.py index 40d30f3..434a6d3 100644 --- a/tests/unit/test_gsm_registries.py +++ b/tests/unit/test_gsm_registries.py @@ -16,7 +16,7 @@ from kilostar.core.global_state_machine import global_state_machine as gsm_modul @pytest.fixture def gsm_instance(monkeypatch): - GSMClass = gsm_module.GlobalStateMachine.__ray_actor_class__ + GSMClass = gsm_module.GlobalStateMachine obj = GSMClass.__new__(GSMClass) obj._mcp_servers = {} obj._tool_configs = {} @@ -50,30 +50,33 @@ def gsm_instance(monkeypatch): async def test_add_mcp_server(gsm_instance): obj = gsm_instance saved = {"server_id": "fs", "name": "fs", "transport": "stdio"} - obj.postgres_database.upsert_mcp_server.remote = AsyncMock(return_value=saved) + obj.postgres_database.upsert_mcp_server = AsyncMock(return_value=saved) ok = await obj.add_mcp_server("fs", {"name": "fs", "transport": "stdio"}) assert ok is True assert obj._mcp_servers["fs"] == saved -def test_get_mcp_server_configs_returns_copy(gsm_instance): +@pytest.mark.asyncio +async def test_get_mcp_server_configs_returns_copy(gsm_instance): obj = gsm_instance obj._mcp_servers["fs"] = {"name": "fs", "transport": "stdio"} - res1 = obj.get_mcp_server_configs() + res1 = await obj.get_mcp_server_configs() res1["fs"] = {"mutated": True} - res2 = obj.get_mcp_server_configs() + res2 = await obj.get_mcp_server_configs() assert res2["fs"]["name"] == "fs" -def test_get_mcp_server_returns_none_when_missing(gsm_instance): - assert gsm_instance.get_mcp_server("nope") is None +@pytest.mark.asyncio +async def test_get_mcp_server_returns_none_when_missing(gsm_instance): + assert await gsm_instance.get_mcp_server("nope") is None -def test_list_mcp_servers_includes_server_id(gsm_instance): +@pytest.mark.asyncio +async def test_list_mcp_servers_includes_server_id(gsm_instance): obj = gsm_instance obj._mcp_servers["fs"] = {"name": "fs", "transport": "stdio"} - listed = obj.list_mcp_servers() + listed = await obj.list_mcp_servers() assert listed[0]["server_id"] == "fs" assert listed[0]["name"] == "fs" @@ -82,7 +85,7 @@ def test_list_mcp_servers_includes_server_id(gsm_instance): async def test_delete_mcp_server(gsm_instance): obj = gsm_instance obj._mcp_servers["fs"] = {"name": "fs"} - obj.postgres_database.delete_mcp_server_db.remote = AsyncMock(return_value=True) + obj.postgres_database.delete_mcp_server_db = AsyncMock(return_value=True) assert await obj.delete_mcp_server("fs") is True assert "fs" not in obj._mcp_servers @@ -91,7 +94,7 @@ async def test_delete_mcp_server(gsm_instance): @pytest.mark.asyncio async def test_delete_unknown_mcp_server(gsm_instance): obj = gsm_instance - obj.postgres_database.delete_mcp_server_db.remote = AsyncMock(return_value=False) + obj.postgres_database.delete_mcp_server_db = AsyncMock(return_value=False) assert await obj.delete_mcp_server("nope") is False @@ -101,46 +104,49 @@ async def test_delete_unknown_mcp_server(gsm_instance): @pytest.mark.asyncio async def test_set_and_get_tool_config(gsm_instance): obj = gsm_instance - obj.postgres_database.upsert_tool_config.remote = AsyncMock( + obj.postgres_database.upsert_tool_config = AsyncMock( return_value={"tool_name": "tavily_search", "config": {"api_key": "xxx"}} ) await obj.set_tool_config("tavily_search", {"api_key": "xxx"}) - assert obj.get_tool_config("tavily_search") == {"api_key": "xxx"} + assert await obj.get_tool_config("tavily_search") == {"api_key": "xxx"} -def test_get_unknown_tool_config_returns_empty(gsm_instance): - assert gsm_instance.get_tool_config("not_exist") == {} +@pytest.mark.asyncio +async def test_get_unknown_tool_config_returns_empty(gsm_instance): + assert await gsm_instance.get_tool_config("not_exist") == {} -def test_get_tool_config_is_isolated_copy(gsm_instance): +@pytest.mark.asyncio +async def test_get_tool_config_is_isolated_copy(gsm_instance): obj = gsm_instance obj._tool_configs["tavily_search"] = {"api_key": "xxx"} - snapshot = obj.get_tool_config("tavily_search") + snapshot = await obj.get_tool_config("tavily_search") snapshot["api_key"] = "changed" - assert obj.get_tool_config("tavily_search") == {"api_key": "xxx"} + assert await obj.get_tool_config("tavily_search") == {"api_key": "xxx"} @pytest.mark.asyncio async def test_delete_tool_config(gsm_instance): obj = gsm_instance obj._tool_configs["tavily_search"] = {"api_key": "xxx"} - obj.postgres_database.delete_tool_config_db.remote = AsyncMock(return_value=True) + obj.postgres_database.delete_tool_config_db = AsyncMock(return_value=True) assert await obj.delete_tool_config("tavily_search") is True - assert obj.get_tool_config("tavily_search") == {} + assert await obj.get_tool_config("tavily_search") == {} @pytest.mark.asyncio async def test_delete_unknown_tool_config(gsm_instance): obj = gsm_instance - obj.postgres_database.delete_tool_config_db.remote = AsyncMock(return_value=False) + obj.postgres_database.delete_tool_config_db = AsyncMock(return_value=False) assert await obj.delete_tool_config("not_exist") is False -def test_list_tool_configs(gsm_instance): +@pytest.mark.asyncio +async def test_list_tool_configs(gsm_instance): obj = gsm_instance obj._tool_configs["tavily_search"] = {"api_key": "xxx"} obj._tool_configs["notion"] = {"token": "yyy"} - raw = obj.list_tool_configs() + raw = await obj.list_tool_configs() assert raw["tavily_search"] == {"api_key": "xxx"} assert raw["notion"] == {"token": "yyy"} @@ -152,7 +158,7 @@ def test_list_tool_configs(gsm_instance): async def test_add_custom_toolset_success(gsm_instance): obj = gsm_instance saved = {"toolset_id": "t1", "name": "my-set", "tools": ["tp_a", "tp_b"]} - obj.postgres_database.upsert_custom_toolset.remote = AsyncMock(return_value=saved) + obj.postgres_database.upsert_custom_toolset = AsyncMock(return_value=saved) result = await obj.add_custom_toolset( toolset_id="t1", name="my-set", tools=["tp_a", "tp_b"] @@ -171,24 +177,26 @@ async def test_add_custom_toolset_rejects_system_tools(gsm_instance): ) -def test_list_custom_toolsets(gsm_instance): +@pytest.mark.asyncio +async def test_list_custom_toolsets(gsm_instance): obj = gsm_instance obj._custom_toolsets["t1"] = {"toolset_id": "t1", "name": "a", "tools": []} - assert len(obj.list_custom_toolsets()) == 1 + assert len(await obj.list_custom_toolsets()) == 1 -def test_get_custom_toolset(gsm_instance): +@pytest.mark.asyncio +async def test_get_custom_toolset(gsm_instance): obj = gsm_instance obj._custom_toolsets["t1"] = {"toolset_id": "t1", "name": "a"} - assert obj.get_custom_toolset("t1")["name"] == "a" - assert obj.get_custom_toolset("nope") is None + assert (await obj.get_custom_toolset("t1"))["name"] == "a" + assert await obj.get_custom_toolset("nope") is None @pytest.mark.asyncio async def test_delete_custom_toolset(gsm_instance): obj = gsm_instance obj._custom_toolsets["t1"] = {"toolset_id": "t1"} - obj.postgres_database.delete_custom_toolset.remote = AsyncMock(return_value=True) + obj.postgres_database.delete_custom_toolset = AsyncMock(return_value=True) assert await obj.delete_custom_toolset("t1") is True assert "t1" not in obj._custom_toolsets obj._global_tool_manager.rebuild_custom_toolsets.assert_called() diff --git a/tests/unit/test_gsm_snapshot.py b/tests/unit/test_gsm_snapshot.py index 9bfc800..5e311d0 100644 --- a/tests/unit/test_gsm_snapshot.py +++ b/tests/unit/test_gsm_snapshot.py @@ -1,26 +1,23 @@ -"""GSM 配置快照(Object Store 读路径)相关测试。 +"""GSM 配置快照(进程内读路径)相关测试。 + +去 Ray 后,快照不再进 Ray Object Store:``_publish_snapshot`` 直接在进程内构建 +并持有一个不可变 ``GSMSnapshot`` 对象,``current_config_ref`` 返回 ``(version, snapshot)``。 +读端 ``fetch_snapshot`` 用版本号做进程内缓存失效。 主要验证: -- ``GSMSnapshot`` 数据类可被 cloudpickle 序列化(ray.put 的隐式约束) - ``_build_snapshot`` 正确从 6 类内存状态打包配置 -- ``_publish_snapshot`` 让 version 单调递增并刷新 ObjectRef -- 写入路径(add_individual / set_tool_config / 等)会自动发布新快照 +- ``_publish_snapshot`` 让 version 单调递增并刷新快照对象 +- 写入路径(add_individual / add_provider_wrap / 等)会自动发布新快照 - ``fetch_snapshot`` 客户端:版本号一致时走本地缓存,不一致时重拉 """ from __future__ import annotations -import asyncio -import pickle -from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock import pytest -# cloudpickle 是 ray 的传递依赖,不直接列在 pyproject 里 —— 通过 ray._private 拿 -from ray import cloudpickle - from kilostar.core.global_state_machine.gsm_snapshot import ( GSMSnapshot, fetch_snapshot, @@ -28,61 +25,14 @@ from kilostar.core.global_state_machine.gsm_snapshot import ( ) -def test_empty_snapshot_can_cloudpickle_roundtrip(): - """空 snapshot 序列化反序列化语义不变(ray.put 的最低约束)。""" - snap = GSMSnapshot() - blob = cloudpickle.dumps(snap) - restored: GSMSnapshot = cloudpickle.loads(blob) - assert restored.version == 0 - assert restored.providers == {} - assert restored.individuals == {} - - -def test_snapshot_with_real_data_roundtrip(): - """带真实 Provider + 函数引用 + dict 数据的 snapshot 也能 round-trip。""" - from kilostar.core.global_state_machine.model_provider.base_provider import ( - Provider, - ) - - def _sample_tool(query: str) -> str: - return f"echo:{query}" - - snap = GSMSnapshot( - version=42, - providers={ - "p1": Provider( - provider_title="p1", - provider_url="http://x", - provider_apikey="sk-x", - provider_models=["gpt-4o"], - provider_type="openai", - ), - }, - individuals={"agent-a": {"agent_id": "agent-a", "model_id": "gpt-4o"}}, - tool_funcs={"echo": _sample_tool}, - ) - blob = cloudpickle.dumps(snap) - restored: GSMSnapshot = cloudpickle.loads(blob) - assert restored.version == 42 - assert restored.providers["p1"].provider_title == "p1" - assert restored.individuals["agent-a"]["model_id"] == "gpt-4o" - # 模块级函数 cloudpickle 后仍可调用 - # 注意:此处函数是测试模块的局部,cloudpickle 会把字节码一并序列化 - assert restored.tool_funcs["echo"]("hi") == "echo:hi" - - -# ─── GSM actor 集成(绕过 @ray.remote 直接构造) ──────────────────── +# ─── GSM 集成(直接构造 plain 对象) ──────────────────────────────── @pytest.fixture -def gsm_instance(monkeypatch): +def gsm_instance(): from kilostar.core.global_state_machine.global_state_machine import ( GlobalStateMachine, ) - - cls = GlobalStateMachine.__ray_actor_class__ - obj = cls.__new__(cls) - # 手动还原 __init__ 副作用 from kilostar.core.global_state_machine.individual_manager import ( GlobalIndividualManager, ) @@ -90,6 +40,7 @@ def gsm_instance(monkeypatch): from kilostar.core.global_state_machine.skill_manager import GlobalSkillManager from kilostar.core.global_state_machine.tool_manager import GlobalToolManager + obj = GlobalStateMachine.__new__(GlobalStateMachine) obj._global_provider_manager = ProviderManager(postgres=None) obj._global_tool_manager = GlobalToolManager() obj._global_skill_manager = GlobalSkillManager() @@ -100,18 +51,6 @@ def gsm_instance(monkeypatch): obj._config_version = 0 obj._current_ref = None obj.postgres_database = MagicMock() - - # ray.put 在测试沙箱里因 psutil PID 检查失败,mock 成"返回一个 sentinel ref" - # 我们关心的是 _publish_snapshot 的语义流,不是真把对象塞进 plasma - import kilostar.core.global_state_machine.global_state_machine as gsm_mod - - counter = {"n": 0} - - def _fake_put(snapshot): - counter["n"] += 1 - return f"fake-ref-{counter['n']}" - - monkeypatch.setattr(gsm_mod.ray, "put", _fake_put) return obj @@ -145,7 +84,7 @@ def test_build_snapshot_picks_up_all_six_categories(gsm_instance): def test_build_snapshot_exposes_system_tools_by_scope(gsm_instance): """系统工具按 scope 分桶的工具名清单要随快照发布出去(客户端重建 toolset 用)。""" tm = gsm_instance._global_tool_manager - # 模拟 tool_manager 内部状态:default scope 有 file_reader,control_node 有 approval + def _f1(): return "f1" @@ -159,7 +98,6 @@ def test_build_snapshot_exposes_system_tools_by_scope(gsm_instance): snap = gsm_instance._build_snapshot() assert snap.system_tools_by_scope.get("default") == ["file_reader"] assert snap.system_tools_by_scope.get("control_node") == ["approval"] - # tool_funcs 拍平后两者都应存在 assert set(snap.tool_funcs.keys()) == {"file_reader", "approval"} @@ -175,15 +113,15 @@ def test_publish_snapshot_increments_version(gsm_instance): gsm_instance._publish_snapshot() assert gsm_instance._config_version == 2 - assert gsm_instance._current_ref is not ref1 # 新 put 应是新 ref + assert gsm_instance._current_ref is not ref1 # 重新构建应是新对象 @pytest.mark.asyncio async def test_current_config_ref_lazy_publishes_when_empty(gsm_instance): """从未发布过快照时,current_config_ref 应自动发布一次而不是返回 None。""" - version, ref = await gsm_instance.current_config_ref() + version, snap = await gsm_instance.current_config_ref() assert version == 1 - assert ref is not None + assert snap is not None @pytest.mark.asyncio @@ -197,7 +135,7 @@ async def test_current_version_is_lightweight(gsm_instance): async def test_add_individual_publishes_new_snapshot(gsm_instance): """写入路径 add_individual 应自动 +1 version。""" before = gsm_instance._config_version - gsm_instance.add_individual("agent-x", {"model_id": "gpt-4o"}) + await gsm_instance.add_individual("agent-x", {"model_id": "gpt-4o"}) after = gsm_instance._config_version assert after == before + 1 @@ -220,8 +158,7 @@ async def test_add_provider_wrap_publishes_new_snapshot(gsm_instance): gsm_instance._global_provider_manager.provider_mapper[ "openai" ].create_provider = AsyncMock(return_value=fake_provider) - gsm_instance.postgres_database.add_provider_db = MagicMock() - gsm_instance.postgres_database.add_provider_db.remote = AsyncMock() + gsm_instance.postgres_database.add_provider_db = AsyncMock() before = gsm_instance._config_version await gsm_instance.add_provider_wrap( @@ -240,69 +177,53 @@ async def test_add_provider_wrap_publishes_new_snapshot(gsm_instance): @pytest.mark.asyncio async def test_fetch_snapshot_uses_local_cache_when_version_matches(): - """模拟 GSM actor,验证版本号一致时不走 ray.get。""" + """版本号一致时不再调 current_config_ref,直接返回本地缓存快照。""" reset_local_cache() snap = GSMSnapshot(version=5, providers={"p": MagicMock()}) - # mock GSM handle:第一次 fetch 全走,第二次只 current_version fake_gsm = MagicMock() - fake_gsm.current_version = MagicMock() - fake_gsm.current_version.remote = AsyncMock(return_value=5) - fake_gsm.current_config_ref = MagicMock() + fake_gsm.current_version = AsyncMock(return_value=5) + fake_gsm.current_config_ref = AsyncMock( + side_effect=AssertionError("不应触发:缓存版本一致时不应调 current_config_ref") + ) - # 提前把缓存预热成 v5(模拟之前已经 fetch 过) from kilostar.core.global_state_machine import gsm_snapshot as snap_mod snap_mod._local_cache["version"] = 5 snap_mod._local_cache["snapshot"] = snap - # 不 mock current_config_ref —— 如果它被调用了,AttributeError 会让测试失败 - fake_gsm.current_config_ref.remote = AsyncMock( - side_effect=AssertionError("不应触发:缓存版本一致时不应调 current_config_ref") - ) - result = await fetch_snapshot(gsm_actor=fake_gsm) assert result is snap - fake_gsm.current_version.remote.assert_awaited_once() + fake_gsm.current_version.assert_awaited_once() @pytest.mark.asyncio -async def test_fetch_snapshot_refetches_when_version_changes(monkeypatch): - """版本号变了应重新 ray.get 拉新 snapshot。""" +async def test_fetch_snapshot_refetches_when_version_changes(): + """版本号变了应重新拉取新 snapshot 并更新缓存。""" reset_local_cache() new_snap = GSMSnapshot(version=10) fake_gsm = MagicMock() - fake_gsm.current_version = MagicMock() - fake_gsm.current_version.remote = AsyncMock(return_value=10) - fake_gsm.current_config_ref = MagicMock() - fake_gsm.current_config_ref.remote = AsyncMock(return_value=(10, "fake-ref")) - - # mock ray.get 让它直接返回我们准备的 snap - import kilostar.core.global_state_machine.gsm_snapshot as snap_mod - - monkeypatch.setattr(snap_mod.ray, "get", lambda ref: new_snap) + fake_gsm.current_version = AsyncMock(return_value=10) + fake_gsm.current_config_ref = AsyncMock(return_value=(10, new_snap)) result = await fetch_snapshot(gsm_actor=fake_gsm) assert result is new_snap - fake_gsm.current_config_ref.remote.assert_awaited_once() - # 缓存应已更新到 v10 + fake_gsm.current_config_ref.assert_awaited_once() + + from kilostar.core.global_state_machine import gsm_snapshot as snap_mod + assert snap_mod._local_cache["version"] == 10 @pytest.mark.asyncio -async def test_fetch_snapshot_use_cache_false_skips_cache(monkeypatch): +async def test_fetch_snapshot_use_cache_false_skips_cache(): """``use_cache=False`` 直接走 current_config_ref,不读本地缓存。""" reset_local_cache() fresh = GSMSnapshot(version=1) fake_gsm = MagicMock() - fake_gsm.current_config_ref = MagicMock() - fake_gsm.current_config_ref.remote = AsyncMock(return_value=(1, "ref")) - - import kilostar.core.global_state_machine.gsm_snapshot as snap_mod - - monkeypatch.setattr(snap_mod.ray, "get", lambda ref: fresh) + fake_gsm.current_config_ref = AsyncMock(return_value=(1, fresh)) result = await fetch_snapshot(gsm_actor=fake_gsm, use_cache=False) assert result is fresh @@ -329,7 +250,10 @@ def test_build_tools_for_scope_assembles_system_and_custom(): snap = GSMSnapshot( all_funcs={"sys_default": _sys_default, "sys_scope": _sys_scope, "tp_a": _tp_a}, custom_toolsets={ - "system_basic": {"toolset_id": "system_basic", "tools": ["sys_default", "sys_scope"]}, + "system_basic": { + "toolset_id": "system_basic", + "tools": ["sys_default", "sys_scope"], + }, "grp": {"toolset_id": "grp", "tools": ["tp_a"]}, }, ) @@ -345,8 +269,5 @@ def test_build_tools_for_scope_skips_empty_buckets(): build_tools_for_scope, ) - snap = GSMSnapshot( - all_funcs={}, - custom_toolsets={}, - ) + snap = GSMSnapshot(all_funcs={}, custom_toolsets={}) assert build_tools_for_scope(snap, "control_node") == [] diff --git a/tests/unit/test_individual_nodes.py b/tests/unit/test_individual_nodes.py index 554115b..2897d81 100644 --- a/tests/unit/test_individual_nodes.py +++ b/tests/unit/test_individual_nodes.py @@ -20,7 +20,7 @@ def regulatory_instance(): from kilostar.core.individual.regulatory_node.regulatory_node import ( RegulatoryNode, ) - cls = RegulatoryNode.__ray_actor_class__ + cls = RegulatoryNode obj = cls.__new__(cls) from kilostar.utils.logger import get_logger obj.logger = get_logger("regulatory_node") @@ -86,7 +86,7 @@ def consciousness_instance(): from kilostar.core.individual.consciousness_node.consciousness_node import ( ConsciousnessNode, ) - cls = ConsciousnessNode.__ray_actor_class__ + cls = ConsciousnessNode obj = cls.__new__(cls) from kilostar.utils.logger import get_logger obj.logger = get_logger("consciousness_node") diff --git a/tests/unit/test_plugin_runtime.py b/tests/unit/test_plugin_runtime.py index 094b253..d4b8b33 100644 --- a/tests/unit/test_plugin_runtime.py +++ b/tests/unit/test_plugin_runtime.py @@ -141,17 +141,11 @@ async def test_on_first_install_default_is_noop(): async def test_install_marker_drives_on_first_install(tmp_path, monkeypatch): """首次装载触发 on_first_install + 写 marker;二次装载不再触发。""" from kilostar.utils import settings as _settings_mod - from kilostar.utils import ray_compat - monkeypatch.setattr(ray_compat, "_STANDALONE", True) monkeypatch.setenv("KILOSTAR_PLUGIN_DIR", str(tmp_path)) _settings_mod.get_settings.cache_clear() - # 强制走 standalone 分支,重新 import plugin_manager 以 re-decorate - import importlib - import kilostar.plugin_runtime.plugin_manager as pm_mod - importlib.reload(pm_mod) - GlobalPluginManager = pm_mod.GlobalPluginManager + from kilostar.plugin_runtime.plugin_manager import GlobalPluginManager plugin_dir = tmp_path / "demo_plug" (plugin_dir / "core").mkdir(parents=True) diff --git a/tests/unit/test_provider_manager.py b/tests/unit/test_provider_manager.py index be6d4b9..91181d0 100644 --- a/tests/unit/test_provider_manager.py +++ b/tests/unit/test_provider_manager.py @@ -30,8 +30,7 @@ async def test_add_provider_happy_path_writes_register_and_db(): pm.provider_mapper["openai"].create_provider = AsyncMock(return_value=fake_provider) postgres = MagicMock() - postgres.add_provider_db = MagicMock() - postgres.add_provider_db.remote = AsyncMock(return_value=None) + postgres.add_provider_db = AsyncMock(return_value=None) await pm.add_provider( provider_type="openai", @@ -44,8 +43,8 @@ async def test_add_provider_happy_path_writes_register_and_db(): assert "my-openai" in pm.provider_register assert pm.provider_register["my-openai"] is fake_provider - postgres.add_provider_db.remote.assert_awaited_once() - kwargs = postgres.add_provider_db.remote.await_args.kwargs + postgres.add_provider_db.assert_awaited_once() + kwargs = postgres.add_provider_db.await_args.kwargs assert kwargs["provider_title"] == "my-openai" assert kwargs["provider_apikey"] == "sk-xxx" assert kwargs["provider_models"] == ["gpt-4o"] @@ -55,8 +54,7 @@ async def test_add_provider_happy_path_writes_register_and_db(): async def test_add_provider_unknown_type_returns_none(caplog): pm = ProviderManager(postgres=None) postgres = MagicMock() - postgres.add_provider_db = MagicMock() - postgres.add_provider_db.remote = AsyncMock() + postgres.add_provider_db = AsyncMock() result = await pm.add_provider( provider_type="not_supported", @@ -69,7 +67,7 @@ async def test_add_provider_unknown_type_returns_none(caplog): assert result is None assert "x" not in pm.provider_register - postgres.add_provider_db.remote.assert_not_awaited() + postgres.add_provider_db.assert_not_awaited() @pytest.mark.asyncio @@ -85,8 +83,7 @@ async def test_add_provider_network_error_raises_retryable(): ) postgres = MagicMock() - postgres.add_provider_db = MagicMock() - postgres.add_provider_db.remote = AsyncMock() + postgres.add_provider_db = AsyncMock() with pytest.raises(RetryableError): await pm.add_provider( diff --git a/tests/unit/test_ray_compat.py b/tests/unit/test_ray_compat.py deleted file mode 100644 index aa1cf67..0000000 --- a/tests/unit/test_ray_compat.py +++ /dev/null @@ -1,106 +0,0 @@ -"""ray_compat 适配层单元测试。 - -验证 StandaloneProxy / _MethodProxy / actor_class / remote_task -在单机模式下的行为是否正确模拟了 Ray Actor Handle 的 .remote() 接口。 -""" - -import asyncio -import pytest - -from kilostar.utils import ray_compat -from kilostar.utils.ray_compat import StandaloneProxy, _MethodProxy - - -class TestMethodProxy: - def test_sync_method(self): - def add(a, b): - return a + b - - proxy = _MethodProxy(add) - result = asyncio.get_event_loop().run_until_complete(proxy.remote(2, 3)) - assert result == 5 - - def test_async_method(self): - async def async_add(a, b): - return a + b - - proxy = _MethodProxy(async_add) - result = asyncio.get_event_loop().run_until_complete(proxy.remote(4, 6)) - assert result == 10 - - -class TestStandaloneProxy: - def test_method_call(self): - class FakeActor: - def greet(self, name): - return f"hello {name}" - - proxy = StandaloneProxy(FakeActor()) - future = proxy.greet.remote("world") - result = asyncio.get_event_loop().run_until_complete(future) - assert result == "hello world" - - def test_async_method_call(self): - class FakeActor: - async def compute(self, x): - return x * 2 - - proxy = StandaloneProxy(FakeActor()) - future = proxy.compute.remote(7) - result = asyncio.get_event_loop().run_until_complete(future) - assert result == 14 - - def test_attribute_access(self): - class FakeActor: - def __init__(self): - self.name = "test" - - proxy = StandaloneProxy(FakeActor()) - assert proxy.name == "test" - - -class TestActorClass: - def test_standalone_returns_class_unchanged(self, monkeypatch): - monkeypatch.setattr(ray_compat, "_STANDALONE", True) - - @ray_compat.actor_class - class MyActor: - def do_work(self): - return 42 - - instance = MyActor() - assert instance.do_work() == 42 - - def test_standalone_class_is_plain_python(self, monkeypatch): - monkeypatch.setattr(ray_compat, "_STANDALONE", True) - - @ray_compat.actor_class - class MyActor: - pass - - assert not hasattr(MyActor, "remote") - assert not hasattr(MyActor, "options") - - -class TestRemoteTask: - def test_sync_task(self, monkeypatch): - monkeypatch.setattr(ray_compat, "_STANDALONE", True) - - @ray_compat.remote_task - def multiply(a, b): - return a * b - - future = multiply.remote(3, 4) - result = asyncio.get_event_loop().run_until_complete(future) - assert result == 12 - - def test_async_task(self, monkeypatch): - monkeypatch.setattr(ray_compat, "_STANDALONE", True) - - @ray_compat.remote_task - async def async_multiply(a, b): - return a * b - - future = async_multiply.remote(5, 6) - result = asyncio.get_event_loop().run_until_complete(future) - assert result == 30 diff --git a/tests/unit/test_request_context.py b/tests/unit/test_request_context.py index d33bcff..4922630 100644 --- a/tests/unit/test_request_context.py +++ b/tests/unit/test_request_context.py @@ -4,7 +4,7 @@ - ``request_id`` / ``trace_id`` 默认空、bind 后可读、reset 后还原 - ``request_id_scope`` / ``trace_id_scope`` 上下文管理器 -- ``snapshot`` / ``apply_snapshot`` 跨边界透传 +- ``snapshot`` 快照当前上下文的 ID - logger 切面:``contextvars`` 中的值会自动写入 ``record["extra"]`` """ @@ -62,21 +62,6 @@ def test_snapshot_returns_current_ids(): assert snap == {"request_id": "r1", "trace_id": "t1"} -def test_apply_snapshot_restores_after_exit(): - snap = {"request_id": "r2", "trace_id": "t2"} - with rc.apply_snapshot(snap): - assert rc.get_request_id() == "r2" - assert rc.get_trace_id() == "t2" - assert rc.get_request_id() == "" - assert rc.get_trace_id() == "" - - -def test_apply_snapshot_handles_none(): - """传 None 应是 no-op,不报错。""" - with rc.apply_snapshot(None): - assert rc.get_request_id() == "" - - @pytest.mark.asyncio async def test_contextvars_isolated_between_concurrent_tasks(): """两个并发的 asyncio task 各自的 trace_id 不应互相串味。""" diff --git a/tests/unit/test_utils_actor.py b/tests/unit/test_utils_actor.py new file mode 100644 index 0000000..d43aad8 --- /dev/null +++ b/tests/unit/test_utils_actor.py @@ -0,0 +1,94 @@ +"""进程内 actor 运行时(``kilostar.utils.actor``)单元测试。 + +覆盖:注册/寻址/注销;句柄对同步与异步方法都返回 awaitable;便捷 getter; +未注册抛 KeyError。 +""" + +from __future__ import annotations + +import pytest + +from kilostar.utils import actor as actor_mod +from kilostar.utils.actor import ( + clear_actors, + get_actor, + get_gsm, + get_gwm, + get_postgres, + register_actor, + unregister_actor, +) + + +@pytest.fixture(autouse=True) +def _clean_registry(): + clear_actors() + yield + clear_actors() + + +class _Sample: + def __init__(self) -> None: + self.value = 7 + + def sync_method(self, x): + return x + 1 + + async def async_method(self, x): + return x * 2 + + +@pytest.mark.asyncio +async def test_sync_method_is_awaitable_through_handle(): + register_actor("sample", _Sample()) + handle = get_actor("sample") + assert await handle.sync_method(41) == 42 + + +@pytest.mark.asyncio +async def test_async_method_forwards_through_handle(): + register_actor("sample", _Sample()) + handle = get_actor("sample") + assert await handle.async_method(21) == 42 + + +def test_non_callable_attribute_passthrough(): + register_actor("sample", _Sample()) + assert get_actor("sample").value == 7 + + +def test_get_unregistered_raises_keyerror(): + with pytest.raises(KeyError): + get_actor("does_not_exist") + + +def test_unregister_then_missing(): + register_actor("sample", _Sample()) + assert get_actor("sample") is not None + unregister_actor("sample") + with pytest.raises(KeyError): + get_actor("sample") + # 幂等:再次注销不报错 + unregister_actor("sample") + + +def test_convenience_getters_resolve_named_actors(): + register_actor("postgres_database", _Sample()) + register_actor("global_state_machine", _Sample()) + register_actor("global_workflow_manager", _Sample()) + assert get_postgres() is get_actor("postgres_database") + assert get_gsm() is get_actor("global_state_machine") + assert get_gwm() is get_actor("global_workflow_manager") + + +@pytest.mark.asyncio +async def test_reregister_overwrites_underlying_instance(): + first = _Sample() + first.value = 100 + register_actor("sample", first) + assert get_actor("sample").value == 100 + + second = _Sample() + second.value = 200 + register_actor("sample", second) + assert get_actor("sample").value == 200 diff --git a/tests/unit/test_utils_get_tool.py b/tests/unit/test_utils_get_tool.py deleted file mode 100644 index f9479a5..0000000 --- a/tests/unit/test_utils_get_tool.py +++ /dev/null @@ -1,44 +0,0 @@ -"""``utils.get_tool`` 在真实仓库目录上的加载行为。""" - -from kilostar.utils import get_tool -from kilostar.utils.get_tool import ( - _get_tool_func, - del_tool_cache, - load_tools_from_list, -) - - -def setup_function(_func): - """每个测试前清空模块级缓存,避免相互影响。""" - get_tool._tool_cache.clear() - - -def test_load_existing_tool_via_load_tools_from_list(): - tools = load_tools_from_list(["file_reader"]) - assert len(tools) == 1 - assert tools[0].__name__ == "file_reader" - - -def test_loader_caches_function(): - f1 = _get_tool_func("file_reader") - f2 = _get_tool_func("file_reader") - assert f1 is f2 - assert "file_reader" in get_tool._tool_cache - - -def test_del_tool_cache_removes_entry(): - _get_tool_func("file_reader") - assert "file_reader" in get_tool._tool_cache - del_tool_cache("file_reader") - assert "file_reader" not in get_tool._tool_cache - - -def test_load_unknown_tool_returns_none_and_keeps_others(): - tools = load_tools_from_list(["file_reader", "definitely_not_exist"]) - assert len(tools) == 1 - assert tools[0].__name__ == "file_reader" - - -def test_load_tools_from_list_handles_none_and_empty(): - assert load_tools_from_list(None) == [] - assert load_tools_from_list([]) == [] diff --git a/tests/unit/test_utils_ray_hook.py b/tests/unit/test_utils_ray_hook.py deleted file mode 100644 index 3fbbd55..0000000 --- a/tests/unit/test_utils_ray_hook.py +++ /dev/null @@ -1,107 +0,0 @@ -"""``ray_hook`` 中纯逻辑容器与 actor 句柄缓存的行为。""" - -from unittest.mock import MagicMock - -import pytest - -from kilostar.utils.ray_hook import ActorList - - -def test_actor_list_attribute_set_get_delete(): - actors = ActorList() - actors.foo = "bar" - assert actors.foo == "bar" - del actors.foo - with pytest.raises(AttributeError): - _ = actors.foo - - -def test_actor_list_missing_raises_attribute_error(): - actors = ActorList() - with pytest.raises(AttributeError): - _ = actors.not_exist - - -def test_actor_list_delete_missing_raises_attribute_error(): - actors = ActorList() - with pytest.raises(AttributeError): - del actors.not_exist - - -def test_ray_actor_hook_uses_fake_registry(fake_actors): - """``ray_actor_hook`` 通过 fake registry 取 actor 并组装成 ActorList。""" - handle = MagicMock() - fake_actors.register("postgres_database", handle) - - from kilostar.utils.ray_hook import ray_actor_hook - - actors = ray_actor_hook("postgres_database") - assert actors.postgres_database is handle - - -def test_ray_actor_hook_unknown_actor_raises(fake_actors): - from kilostar.utils.ray_hook import ray_actor_hook - - with pytest.raises(ValueError): - ray_actor_hook("does_not_exist") - - -def test_wait_for_actor_returns_immediately_when_ready(fake_actors): - """actor 已就绪时 wait_for_actor 立刻返回,不进入轮询等待。""" - handle = MagicMock() - fake_actors.register("postgres_database", handle) - - from kilostar.utils.ray_hook import wait_for_actor - - got = wait_for_actor("postgres_database", timeout=5.0) - assert got is handle - - -def test_wait_for_actor_times_out_with_clear_error(fake_actors): - """超时仍未就绪时抛 TimeoutError,并在 message 里带 actor 名。""" - from kilostar.utils.ray_hook import wait_for_actor - - with pytest.raises(TimeoutError) as exc_info: - wait_for_actor("never_ready", timeout=0.2, interval=0.05) - assert "never_ready" in str(exc_info.value) - - -def test_wait_for_actor_succeeds_after_delayed_registration(fake_actors): - """actor 在第 N 次轮询时才注册,wait_for_actor 应在它就绪后返回。""" - from kilostar.utils.ray_hook import wait_for_actor - - handle = MagicMock() - calls = {"n": 0} - original_get = fake_actors.get - - def delayed_get(name, namespace="kilostar"): - calls["n"] += 1 - if calls["n"] >= 3: - return handle - raise ValueError("not ready yet") - - fake_actors.get = delayed_get - try: - got = wait_for_actor("late_actor", timeout=2.0, interval=0.05) - assert got is handle - assert calls["n"] >= 3 - finally: - fake_actors.get = original_get - - -def test_ray_actor_hook_with_timeout_waits(fake_actors): - """ray_actor_hook(timeout>0) 会走 wait_for_actor 等待路径。""" - from kilostar.utils.ray_hook import ray_actor_hook - - handle = MagicMock() - calls = {"n": 0} - - def delayed_get(name, namespace="kilostar"): - calls["n"] += 1 - if calls["n"] >= 2: - return handle - raise ValueError("not ready yet") - - fake_actors.get = delayed_get - actors = ray_actor_hook("slow_actor", timeout=2.0, interval=0.05) - assert actors.slow_actor is handle diff --git a/tests/unit/test_workflow_engine.py b/tests/unit/test_workflow_engine.py index bbee9bf..49ca744 100644 --- a/tests/unit/test_workflow_engine.py +++ b/tests/unit/test_workflow_engine.py @@ -1,14 +1,14 @@ -"""``ConsciousnessNode.start_workflow_design`` fire workflow ray task 的提交逻辑。 +"""``ConsciousnessNode.start_workflow_design`` fire workflow task 的提交逻辑。 历史上这里有一个常驻的 ``WorkflowRunningEngine`` actor 做中转,现已删除: -workflow 是一次性、有头有尾的执行,更适合直接以 ray task 形式触发。 -本测试保证 ConsciousnessNode 在工作流生成后正确 fire ``run_workflow_task``, +workflow 是一次性、有头有尾的执行,直接以 ``asyncio.create_task`` 触发 +``run_workflow_task``。本测试保证 ConsciousnessNode 在工作流生成后正确 fire, 并通过 ``put_pending`` 推送 SSE 进度(节点端写 pending → API 端 SSE 读 pending)。 """ from __future__ import annotations -from types import SimpleNamespace +import asyncio from unittest.mock import AsyncMock, MagicMock import pytest @@ -25,26 +25,15 @@ def consciousness_instance(): ) from kilostar.utils.logger import get_logger - cls = ConsciousnessNode.__ray_actor_class__ - obj = cls.__new__(cls) + obj = ConsciousnessNode.__new__(ConsciousnessNode) obj.logger = get_logger("consciousness_node") obj.agent = None return obj -class _FakeActorRef: - """模拟 ``ray_actor_hook("name").`` 的链式取属性返回值。""" - - def __init__(self, target): - self._target = target - - def __getattr__(self, item): - return getattr(self._target, item) - - @pytest.mark.asyncio async def test_start_workflow_design_fires_run_workflow_task( - consciousness_instance, monkeypatch + consciousness_instance, fake_actors, monkeypatch ): """快乐路径:working 返回 ForWorkflowEngine,应 fire run_workflow_task 且推送 pending。""" from kilostar.core.individual.consciousness_node.template import ( @@ -62,40 +51,28 @@ async def test_start_workflow_design_fires_run_workflow_task( ) postgres = MagicMock() - postgres.get_all_worker_individual = MagicMock() - postgres.get_all_worker_individual.remote = AsyncMock(return_value=[]) - postgres.update_workflow_status = MagicMock() - postgres.update_workflow_status.remote = AsyncMock() + postgres.get_all_worker_individual = AsyncMock(return_value=[]) + postgres.update_workflow_status = AsyncMock() + fake_actors.register("postgres_database", postgres) pending_writes: list[tuple[str, str]] = [] - gwm = MagicMock() - gwm.put_pending = MagicMock() - gwm.put_pending.remote = AsyncMock( + gwm.put_pending = AsyncMock( side_effect=lambda tid, msg: pending_writes.append((tid, msg)) ) - - def _fake_hook(name): - if name == "postgres_database": - return SimpleNamespace(postgres_database=_FakeActorRef(postgres)) - if name == "global_workflow_manager": - return SimpleNamespace(global_workflow_manager=_FakeActorRef(gwm)) - raise KeyError(name) - - import kilostar.core.individual.consciousness_node.consciousness_node as cmod - - monkeypatch.setattr(cmod, "ray_actor_hook", _fake_hook) + fake_actors.register("global_workflow_manager", gwm) captured: dict = {} - def _fake_task_remote(workflow_dict, trace_id): + async def _fake_task(workflow_dict, trace_id): captured["workflow_dict"] = workflow_dict captured["trace_id"] = trace_id - return MagicMock() - monkeypatch.setattr(engine_module.run_workflow_task, "remote", _fake_task_remote) + monkeypatch.setattr(engine_module, "run_workflow_task", _fake_task) await consciousness_instance.start_workflow_design("trace-123", "do something") + # 让 create_task 调度的协程跑完 + await asyncio.sleep(0) assert captured["trace_id"] == "trace-123" assert captured["workflow_dict"]["title"] == "t" @@ -106,44 +83,33 @@ async def test_start_workflow_design_fires_run_workflow_task( @pytest.mark.asyncio async def test_start_workflow_design_failed_path_marks_failed( - consciousness_instance, monkeypatch + consciousness_instance, fake_actors, monkeypatch ): """working 返回 None / 不匹配类型时应推送失败提示并把 workflow 状态置为 failed。""" consciousness_instance.working = AsyncMock(return_value=None) postgres = MagicMock() - postgres.get_all_worker_individual = MagicMock() - postgres.get_all_worker_individual.remote = AsyncMock(return_value=[]) - postgres.update_workflow_status = MagicMock() - postgres.update_workflow_status.remote = AsyncMock() + postgres.get_all_worker_individual = AsyncMock(return_value=[]) + postgres.update_workflow_status = AsyncMock() + fake_actors.register("postgres_database", postgres) pending_writes: list[tuple[str, str]] = [] gwm = MagicMock() - gwm.put_pending = MagicMock() - gwm.put_pending.remote = AsyncMock( + gwm.put_pending = AsyncMock( side_effect=lambda tid, msg: pending_writes.append((tid, msg)) ) - - def _fake_hook(name): - if name == "postgres_database": - return SimpleNamespace(postgres_database=_FakeActorRef(postgres)) - if name == "global_workflow_manager": - return SimpleNamespace(global_workflow_manager=_FakeActorRef(gwm)) - raise KeyError(name) - - import kilostar.core.individual.consciousness_node.consciousness_node as cmod - - monkeypatch.setattr(cmod, "ray_actor_hook", _fake_hook) + fake_actors.register("global_workflow_manager", gwm) fired: list = [] - monkeypatch.setattr( - engine_module.run_workflow_task, - "remote", - lambda *a, **kw: fired.append((a, kw)), - ) + + async def _fake_task(*a, **kw): + fired.append((a, kw)) + + monkeypatch.setattr(engine_module, "run_workflow_task", _fake_task) await consciousness_instance.start_workflow_design("trace-x", "cmd") + await asyncio.sleep(0) - assert fired == [] # 没有 fire ray task - postgres.update_workflow_status.remote.assert_awaited_with("trace-x", "failed") + assert fired == [] # 没有 fire workflow task + postgres.update_workflow_status.assert_awaited_with("trace-x", "failed") assert any("生成失败" in msg for _, msg in pending_writes)