From 19d1aaf7e9722124b3d4e8b5cd57af6ca42a7442 Mon Sep 17 00:00:00 2001 From: huaiyj <8699003+huaiyj@user.noreply.gitee.com> Date: Mon, 24 Aug 2026 20:53:08 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E5=A2=9E=E5=8A=A0=E4=BC=81=E4=B8=9A?= =?UTF-8?q?=E8=BF=90=E7=BB=B4=E7=9F=A5=E8=AF=86=E5=BA=93?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- ocai-service/ocai/ai/agent_manager.py | 3 + .../ocai/ai/prompts/system_builder.py | 61 ++++++++-- ocai-service/ocai/api/deps.py | 10 +- ocai-service/ocai/api/v2_route.py | 74 ++++++++++++- ocai-service/ocai/main.py | 6 + ocai-service/ocai/repo/interfaces/__init__.py | 2 + .../ocai/repo/interfaces/knowledge.py | 30 +++++ ocai-service/ocai/repo/mysql/__init__.py | 2 + .../ocai/repo/mysql/knowledge_repo.py | 104 ++++++++++++++++++ .../versions/0002_knowledge_entries.py | 40 +++++++ ocai-service/ocai/repo/mysql/models.py | 21 ++++ ocai-service/ocai/schemas/knowledge.py | 29 +++++ ocai-service/ocai/services/__init__.py | 2 + .../ocai/services/knowledge_service.py | 62 +++++++++++ .../sql/0002_knowledge_entries.down.sql | 1 + .../sql/0002_knowledge_entries.up.sql | 15 +++ .../ai/test_enterprise_knowledge_prompt.py | 31 ++++++ .../tests/api/test_knowledge_endpoints.py | 55 +++++++++ ocai-service/tests/repo/test_mysql_repos.py | 46 ++++++++ .../tests/services/test_knowledge_service.py | 54 +++++++++ 20 files changed, 636 insertions(+), 12 deletions(-) create mode 100644 ocai-service/ocai/repo/interfaces/knowledge.py create mode 100644 ocai-service/ocai/repo/mysql/knowledge_repo.py create mode 100644 ocai-service/ocai/repo/mysql/migrations/versions/0002_knowledge_entries.py create mode 100644 ocai-service/ocai/schemas/knowledge.py create mode 100644 ocai-service/ocai/services/knowledge_service.py create mode 100644 ocai-service/sql/0002_knowledge_entries.down.sql create mode 100644 ocai-service/sql/0002_knowledge_entries.up.sql create mode 100644 ocai-service/tests/ai/test_enterprise_knowledge_prompt.py create mode 100644 ocai-service/tests/api/test_knowledge_endpoints.py create mode 100644 ocai-service/tests/services/test_knowledge_service.py diff --git a/ocai-service/ocai/ai/agent_manager.py b/ocai-service/ocai/ai/agent_manager.py index db47dfe..8b3bded 100644 --- a/ocai-service/ocai/ai/agent_manager.py +++ b/ocai-service/ocai/ai/agent_manager.py @@ -111,6 +111,7 @@ class AgentManager: memory_load_fn=None, memory_delete_fn=None, memory_search_fn=None, + knowledge_list_fn=None, ) -> None: """初始化所有 AI 组件 @@ -122,6 +123,7 @@ class AgentManager: memory_load_fn: 记忆加载回调 (user_id) -> List[dict] memory_delete_fn: 记忆删除回调 (memory_id) -> bool memory_search_fn: 记忆搜索回调 (query, user_id, limit) -> List[dict] + knowledge_list_fn: 返回当前用户/项目适用的企业知识条目 """ if self._initialized: return @@ -162,6 +164,7 @@ class AgentManager: skill_manager=self._skill_manager, tool_registry=self._tool_registry, baseline_provider=baseline_provider, + knowledge_provider=knowledge_list_fn, ) # 6. 初始化 Agents diff --git a/ocai-service/ocai/ai/prompts/system_builder.py b/ocai-service/ocai/ai/prompts/system_builder.py index e072eaa..223f410 100644 --- a/ocai-service/ocai/ai/prompts/system_builder.py +++ b/ocai-service/ocai/ai/prompts/system_builder.py @@ -70,6 +70,7 @@ class SystemPromptBuilder: user_rules: str = "", project_rules: str = "", baseline_provider=None, + knowledge_provider=None, ): self._memory_manager = memory_manager self._skill_manager = skill_manager @@ -81,6 +82,7 @@ class SystemPromptBuilder: # 用 (user_id, jwt_token) 触发拉取并按 user_id TTL 缓存。 # 不传则该 section 不出现,其余流程不受影响。 self._baseline_provider = baseline_provider + self._knowledge_provider = knowledge_provider def build( self, @@ -143,12 +145,18 @@ class SystemPromptBuilder: # 6.5 运维基线(请求期按需拉取:用 user_jwt 调 ListProjects, # 按 user_id TTL 缓存;无 JWT 或没配 oc-manager 时自动缺席) - baseline_section = self._build_baseline_section( - user_id=user_id, jwt_token=jwt_token - ) + baseline_section = self._build_baseline_section(user_id=user_id, jwt_token=jwt_token) if baseline_section: sections.append(baseline_section) + knowledge_section = self._build_knowledge_section( + user_id=user_id, + project_id=self._context_value("project_id", preloaded_context, extra_context), + role=self._context_value("role", preloaded_context, extra_context), + ) + if knowledge_section: + sections.append(knowledge_section) + # 7. 用户上下文 (用户信息 + 关联机器 + 预加载资产) user_context = self._build_user_context( user_id=user_id, @@ -206,6 +214,45 @@ class SystemPromptBuilder: return "\n\n".join(sections) + def _build_knowledge_section( + self, + *, + user_id: str | None, + project_id: str | None, + role: str | None, + ) -> str: + if self._knowledge_provider is None: + return "" + try: + entries = self._knowledge_provider( + user_id=user_id, + project_id=project_id, + role=role, + ) + except Exception as exc: + logger.warning("enterprise knowledge lookup failed: %s", exc) + return "" + if not entries: + return "" + parts = ["## 企业运维知识"] + for entry in entries: + title = str(entry.get("title") or "未命名知识") + content = str(entry.get("content") or "").strip() + if content: + parts.extend([f"\n### {title}", content, f"来源: knowledge:{entry.get('id')}"]) + return "\n".join(parts) + + @staticmethod + def _context_value( + key: str, + preloaded_context: dict[str, Any] | None, + extra_context: dict[str, Any] | None, + ) -> str | None: + for source in (preloaded_context, extra_context): + if source and source.get(key) is not None: + return str(source[key]) + return None + def update_with_memory(self, user_message: str, user_id: str | None = None) -> str: """根据用户消息更新系统提示词(注入相关记忆)""" return self.build(user_message, user_id) @@ -249,9 +296,7 @@ class SystemPromptBuilder: return "" try: - baseline = self._baseline_provider.get_baseline( - user_id=user_id, jwt_token=jwt_token - ) + baseline = self._baseline_provider.get_baseline(user_id=user_id, jwt_token=jwt_token) except Exception as exc: # provider 自己已经把异常吞成空 dict,这里再兜一层防御。 logger.warning("baseline_provider.get_baseline crashed: %s", exc) @@ -285,9 +330,7 @@ class SystemPromptBuilder: parts.append( f"\n> 当前用户只参与 1 个项目(id={only.get('id')})。" "调用 `ocm_list_instance_ips` / `ocai_metrics_*` 等工具时," - "**直接使用 ``project_id={pid}``**,不要再问用户。".format( - pid=only.get("id") - ) + "**直接使用 ``project_id={pid}``**,不要再问用户。".format(pid=only.get("id")) ) else: parts.append( diff --git a/ocai-service/ocai/api/deps.py b/ocai-service/ocai/api/deps.py index e977074..44e325e 100644 --- a/ocai-service/ocai/api/deps.py +++ b/ocai-service/ocai/api/deps.py @@ -27,7 +27,7 @@ from fastapi import HTTPException, Request if TYPE_CHECKING: from ocai.schemas.config import ApplicationConfig - from ocai.services import ChatService, UserPreferencesService + from ocai.services import ChatService, KnowledgeService, UserPreferencesService def get_chat_service(request: Request) -> ChatService: @@ -50,6 +50,13 @@ def get_user_preferences_service(request: Request) -> UserPreferencesService: return svc # type: ignore[no-any-return] +def get_knowledge_service(request: Request) -> KnowledgeService: + svc = getattr(request.app.state, "knowledge_service", None) + if svc is None: + raise HTTPException(status_code=500, detail="KnowledgeService not initialized") + return svc # type: ignore[no-any-return] + + def get_app_config(request: Request) -> ApplicationConfig | None: """Resolve the loaded :class:`ApplicationConfig` from ``app.state``. @@ -63,5 +70,6 @@ def get_app_config(request: Request) -> ApplicationConfig | None: __all__ = [ "get_app_config", "get_chat_service", + "get_knowledge_service", "get_user_preferences_service", ] diff --git a/ocai-service/ocai/api/v2_route.py b/ocai-service/ocai/api/v2_route.py index 6280b17..7e1eb21 100644 --- a/ocai-service/ocai/api/v2_route.py +++ b/ocai-service/ocai/api/v2_route.py @@ -38,9 +38,15 @@ from fastapi import APIRouter, Body, Depends, Query, Request from fastapi.responses import JSONResponse, StreamingResponse from starlette.concurrency import iterate_in_threadpool -from ocai.api.deps import get_app_config, get_chat_service, get_user_preferences_service +from ocai.api.deps import ( + get_app_config, + get_chat_service, + get_knowledge_service, + get_user_preferences_service, +) from ocai.permission import UserInfo, get_current_user from ocai.schemas.agent_sse import AgentChatRequest +from ocai.schemas.knowledge import KnowledgeCreateRequest, KnowledgeUpdateRequest from ocai.schemas.v2_ops import ( CancelChatRequest, DeleteSessionRequest, @@ -50,7 +56,7 @@ from ocai.schemas.v2_ops import ( if TYPE_CHECKING: from ocai.schemas.config import ApplicationConfig - from ocai.services import ChatService, UserPreferencesService + from ocai.services import ChatService, KnowledgeService, UserPreferencesService logger = logging.getLogger(__name__) @@ -63,10 +69,74 @@ router = APIRouter(prefix="/api/v2/ocai", tags=["v2"]) ChatSvc = Annotated["ChatService", Depends(get_chat_service)] PrefSvc = Annotated["UserPreferencesService", Depends(get_user_preferences_service)] +KnowledgeSvc = Annotated["KnowledgeService", Depends(get_knowledge_service)] AppConf = Annotated["ApplicationConfig | None", Depends(get_app_config)] CurrentUser = Annotated[UserInfo, Depends(get_current_user)] +@router.get("/knowledge") +async def list_knowledge( + knowledge_service: KnowledgeSvc, + user: CurrentUser, + scope_type: str | None = Query(default=None), + scope_id: str | None = Query(default=None), + enabled: bool | None = Query(default=None), +) -> dict[str, Any]: + return { + "entries": knowledge_service.list( + scope_type=scope_type, + scope_id=scope_id, + enabled=enabled, + created_by=user.login_name, + ) + } + + +@router.post("/knowledge", status_code=201) +async def create_knowledge( + request: KnowledgeCreateRequest, + knowledge_service: KnowledgeSvc, + user: CurrentUser, +) -> dict[str, Any]: + return knowledge_service.create(request, user.login_name) + + +@router.put("/knowledge/{entry_id}") +async def update_knowledge( + entry_id: str, + request: KnowledgeUpdateRequest, + knowledge_service: KnowledgeSvc, + user: CurrentUser, +) -> Any: + try: + result = knowledge_service.update(entry_id, request, user.login_name) + except ValueError as exc: + return JSONResponse( + {"error": {"code": "INVALID_SCOPE", "message": str(exc)}}, + status_code=400, + ) + if result is None: + return JSONResponse( + {"error": {"code": "NOT_FOUND", "message": "知识条目不存在或无权修改"}}, + status_code=404, + ) + return result + + +@router.delete("/knowledge/{entry_id}") +async def delete_knowledge( + entry_id: str, + knowledge_service: KnowledgeSvc, + user: CurrentUser, +) -> Any: + if not knowledge_service.delete(entry_id, user.login_name): + return JSONResponse( + {"error": {"code": "NOT_FOUND", "message": "知识条目不存在或无权删除"}}, + status_code=404, + ) + return {"deleted": True, "id": entry_id} + + # --------------------------------------------------------------------------- # Helpers # --------------------------------------------------------------------------- diff --git a/ocai-service/ocai/main.py b/ocai-service/ocai/main.py index 09e8a99..5d735a0 100644 --- a/ocai-service/ocai/main.py +++ b/ocai-service/ocai/main.py @@ -25,12 +25,14 @@ from ocai.permission import PermissionConfig, setup_permission from ocai.repo.cache_repo import CacheRepository from ocai.repo.mysql import ( MysqlAgentRepository, + MysqlKnowledgeRepository, MysqlUserPreferencesRepo, MysqlV2SessionRepository, ) from ocai.schemas.config import ApplicationConfig from ocai.services import ( ChatService, + KnowledgeService, MessageConverter, UserPreferencesService, ) @@ -197,6 +199,7 @@ def create_app(conf: ApplicationConfig | None = None) -> FastAPI: engine = mysql_db.engine agent_repo = MysqlAgentRepository(engine) + knowledge_repo = MysqlKnowledgeRepository(engine) v2_session_repo = MysqlV2SessionRepository(engine) user_preferences_repo = MysqlUserPreferencesRepo(engine) @@ -216,6 +219,7 @@ def create_app(conf: ApplicationConfig | None = None) -> FastAPI: memory_load_fn=agent_repo.load_agent_memories, memory_delete_fn=agent_repo.delete_agent_memory, memory_search_fn=agent_repo.search_agent_memories, + knowledge_list_fn=knowledge_repo.list_applicable, ) message_converter = MessageConverter() @@ -231,6 +235,7 @@ def create_app(conf: ApplicationConfig | None = None) -> FastAPI: user_preferences_service = UserPreferencesService( user_preferences_repo=user_preferences_repo, ) + knowledge_service = KnowledgeService(knowledge_repo) @asynccontextmanager async def lifespan(_: FastAPI) -> AsyncIterator[None]: @@ -266,6 +271,7 @@ def create_app(conf: ApplicationConfig | None = None) -> FastAPI: app.state.metrics_collector = metrics_collector app.state.chat_service = chat_service app.state.user_preferences_service = user_preferences_service + app.state.knowledge_service = knowledge_service register_exception_handlers(app, logger=logger) diff --git a/ocai-service/ocai/repo/interfaces/__init__.py b/ocai-service/ocai/repo/interfaces/__init__.py index f851564..1c3a28b 100644 --- a/ocai-service/ocai/repo/interfaces/__init__.py +++ b/ocai-service/ocai/repo/interfaces/__init__.py @@ -8,12 +8,14 @@ service code. from ocai.repo.interfaces.agent import AbstractAgentRepository from ocai.repo.interfaces.chat_session import AbstractChatSessionRepository +from ocai.repo.interfaces.knowledge import AbstractKnowledgeRepository from ocai.repo.interfaces.user_preferences import AbstractUserPreferencesRepository from ocai.repo.interfaces.v2_session import AbstractV2SessionRepository __all__ = [ "AbstractAgentRepository", "AbstractChatSessionRepository", + "AbstractKnowledgeRepository", "AbstractUserPreferencesRepository", "AbstractV2SessionRepository", ] diff --git a/ocai-service/ocai/repo/interfaces/knowledge.py b/ocai-service/ocai/repo/interfaces/knowledge.py new file mode 100644 index 0000000..c74bc51 --- /dev/null +++ b/ocai-service/ocai/repo/interfaces/knowledge.py @@ -0,0 +1,30 @@ +from __future__ import annotations + +from abc import ABC, abstractmethod +from typing import Any + + +class AbstractKnowledgeRepository(ABC): + @abstractmethod + def create(self, entry: dict[str, Any]) -> dict[str, Any]: ... + + @abstractmethod + def get(self, entry_id: str) -> dict[str, Any] | None: ... + + @abstractmethod + def list_entries(self, **filters: Any) -> list[dict[str, Any]]: ... + + @abstractmethod + def list_applicable( + self, + *, + user_id: str | None, + project_id: str | None = None, + role: str | None = None, + ) -> list[dict[str, Any]]: ... + + @abstractmethod + def update(self, entry_id: str, values: dict[str, Any]) -> dict[str, Any] | None: ... + + @abstractmethod + def delete(self, entry_id: str) -> bool: ... diff --git a/ocai-service/ocai/repo/mysql/__init__.py b/ocai-service/ocai/repo/mysql/__init__.py index 4ae19b2..b829e19 100644 --- a/ocai-service/ocai/repo/mysql/__init__.py +++ b/ocai-service/ocai/repo/mysql/__init__.py @@ -8,6 +8,7 @@ implementations. from ocai.repo.mysql.agent_repo import MysqlAgentRepository from ocai.repo.mysql.chat_session_repo import MysqlChatSessionRepository from ocai.repo.mysql.engine import MySQLEngine +from ocai.repo.mysql.knowledge_repo import MysqlKnowledgeRepository from ocai.repo.mysql.user_preferences_repo import MysqlUserPreferencesRepo from ocai.repo.mysql.v2_session_repo import MysqlV2SessionRepository @@ -15,6 +16,7 @@ __all__ = [ "MySQLEngine", "MysqlAgentRepository", "MysqlChatSessionRepository", + "MysqlKnowledgeRepository", "MysqlUserPreferencesRepo", "MysqlV2SessionRepository", ] diff --git a/ocai-service/ocai/repo/mysql/knowledge_repo.py b/ocai-service/ocai/repo/mysql/knowledge_repo.py new file mode 100644 index 0000000..7a713f6 --- /dev/null +++ b/ocai-service/ocai/repo/mysql/knowledge_repo.py @@ -0,0 +1,104 @@ +from __future__ import annotations + +from datetime import UTC, datetime +from typing import Any + +from sqlalchemy import and_, case, delete, or_, select + +from ocai.repo.interfaces.knowledge import AbstractKnowledgeRepository +from ocai.repo.mysql.engine import MySQLEngine +from ocai.repo.mysql.models import KnowledgeEntry + + +def _to_dict(row: KnowledgeEntry) -> dict[str, Any]: + return { + "id": row.id, + "title": row.title, + "content": row.content, + "scope_type": row.scope_type, + "scope_id": row.scope_id, + "enabled": bool(row.enabled), + "version": row.version, + "created_by": row.created_by, + "created_at": row.created_at, + "updated_at": row.updated_at, + } + + +class MysqlKnowledgeRepository(AbstractKnowledgeRepository): + def __init__(self, engine: MySQLEngine): + self._engine = engine + + def create(self, entry: dict[str, Any]) -> dict[str, Any]: + now = datetime.now(UTC).replace(tzinfo=None) + row = KnowledgeEntry(**entry, version=1, created_at=now, updated_at=now) + with self._engine.session_scope() as session: + session.add(row) + return _to_dict(row) + + def get(self, entry_id: str) -> dict[str, Any] | None: + with self._engine.session_scope() as session: + row = session.get(KnowledgeEntry, entry_id) + return _to_dict(row) if row else None + + def list_entries(self, **filters: Any) -> list[dict[str, Any]]: + conditions = [] + for name in ("scope_type", "scope_id", "created_by"): + if filters.get(name) is not None: + conditions.append(getattr(KnowledgeEntry, name) == filters[name]) + if filters.get("enabled") is not None: + conditions.append(KnowledgeEntry.enabled == int(bool(filters["enabled"]))) + stmt = select(KnowledgeEntry) + if conditions: + stmt = stmt.where(*conditions) + stmt = stmt.order_by(KnowledgeEntry.updated_at.desc()) + with self._engine.session_scope() as session: + return [_to_dict(row) for row in session.execute(stmt).scalars().all()] + + def list_applicable( + self, + *, + user_id: str | None, + project_id: str | None = None, + role: str | None = None, + ) -> list[dict[str, Any]]: + scopes = [KnowledgeEntry.scope_type == "global"] + if user_id: + scopes.append(and_(KnowledgeEntry.scope_type == "user", KnowledgeEntry.scope_id == user_id)) + if project_id: + scopes.append(and_(KnowledgeEntry.scope_type == "project", KnowledgeEntry.scope_id == project_id)) + if role: + scopes.append(and_(KnowledgeEntry.scope_type == "role", KnowledgeEntry.scope_id == role)) + stmt = ( + select(KnowledgeEntry) + .where(KnowledgeEntry.enabled == 1, or_(*scopes)) + .order_by( + case( + (KnowledgeEntry.scope_type == "global", 0), + (KnowledgeEntry.scope_type == "role", 1), + (KnowledgeEntry.scope_type == "user", 2), + (KnowledgeEntry.scope_type == "project", 3), + else_=0, + ), + KnowledgeEntry.updated_at.desc(), + ) + ) + with self._engine.session_scope() as session: + return [_to_dict(row) for row in session.execute(stmt).scalars().all()] + + def update(self, entry_id: str, values: dict[str, Any]) -> dict[str, Any] | None: + with self._engine.session_scope() as session: + row = session.get(KnowledgeEntry, entry_id) + if row is None: + return None + for key, value in values.items(): + setattr(row, key, value) + row.version += 1 + row.updated_at = datetime.now(UTC).replace(tzinfo=None) + session.flush() + return _to_dict(row) + + def delete(self, entry_id: str) -> bool: + with self._engine.session_scope() as session: + result = session.execute(delete(KnowledgeEntry).where(KnowledgeEntry.id == entry_id)) + return int(getattr(result, "rowcount", 0)) > 0 diff --git a/ocai-service/ocai/repo/mysql/migrations/versions/0002_knowledge_entries.py b/ocai-service/ocai/repo/mysql/migrations/versions/0002_knowledge_entries.py new file mode 100644 index 0000000..f879995 --- /dev/null +++ b/ocai-service/ocai/repo/mysql/migrations/versions/0002_knowledge_entries.py @@ -0,0 +1,40 @@ +"""enterprise knowledge entries + +Revision ID: 0002_knowledge +Revises: 0001_initial +""" + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op + +revision: str = "0002_knowledge" +down_revision: str | None = "0001_initial" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + + +def upgrade() -> None: + op.create_table( + "knowledge_entries", + sa.Column("id", sa.String(length=64), nullable=False), + sa.Column("title", sa.String(length=256), nullable=False), + sa.Column("content", sa.Text(), nullable=False), + sa.Column("scope_type", sa.String(length=32), nullable=False), + sa.Column("scope_id", sa.String(length=128), nullable=True), + sa.Column("enabled", sa.Integer(), nullable=False), + sa.Column("version", sa.Integer(), nullable=False), + sa.Column("created_by", sa.String(length=128), nullable=False), + sa.Column("created_at", sa.DateTime(), nullable=False), + sa.Column("updated_at", sa.DateTime(), nullable=False), + sa.PrimaryKeyConstraint("id"), + ) + op.create_index("idx_knowledge_scope", "knowledge_entries", ["scope_type", "scope_id", "enabled"]) + op.create_index("idx_knowledge_creator", "knowledge_entries", ["created_by", "updated_at"]) + + +def downgrade() -> None: + op.drop_index("idx_knowledge_creator", table_name="knowledge_entries") + op.drop_index("idx_knowledge_scope", table_name="knowledge_entries") + op.drop_table("knowledge_entries") diff --git a/ocai-service/ocai/repo/mysql/models.py b/ocai-service/ocai/repo/mysql/models.py index 415e185..5af9541 100644 --- a/ocai-service/ocai/repo/mysql/models.py +++ b/ocai-service/ocai/repo/mysql/models.py @@ -195,11 +195,32 @@ class UserPreference(Base): updated_at: Mapped[datetime] = mapped_column(DateTime, nullable=False) +class KnowledgeEntry(Base): + __tablename__ = "knowledge_entries" + + id: Mapped[str] = mapped_column(String(64), primary_key=True) + title: Mapped[str] = mapped_column(String(256), nullable=False) + content: Mapped[str] = mapped_column(Text, nullable=False) + scope_type: Mapped[str] = mapped_column(String(32), nullable=False, default="global") + scope_id: Mapped[str | None] = mapped_column(String(128)) + enabled: Mapped[int] = mapped_column(Integer, nullable=False, default=1) + version: Mapped[int] = mapped_column(Integer, nullable=False, default=1) + created_by: Mapped[str] = mapped_column(String(128), nullable=False) + created_at: Mapped[datetime] = mapped_column(DateTime, nullable=False) + updated_at: Mapped[datetime] = mapped_column(DateTime, nullable=False) + + __table_args__ = ( + Index("idx_knowledge_scope", "scope_type", "scope_id", "enabled"), + Index("idx_knowledge_creator", "created_by", "updated_at"), + ) + + __all__ = [ "AgentLog", "AgentMemory", "Base", "ChatSession", + "KnowledgeEntry", "UserPreference", "V2Session", "V2SessionInvolvedAgent", diff --git a/ocai-service/ocai/schemas/knowledge.py b/ocai-service/ocai/schemas/knowledge.py new file mode 100644 index 0000000..ec3c789 --- /dev/null +++ b/ocai-service/ocai/schemas/knowledge.py @@ -0,0 +1,29 @@ +from typing import Literal + +from pydantic import BaseModel, Field, model_validator + +ScopeType = Literal["global", "project", "user", "role"] + + +class KnowledgeCreateRequest(BaseModel): + title: str = Field(min_length=1, max_length=256) + content: str = Field(min_length=1, max_length=50000) + scope_type: ScopeType = "global" + scope_id: str | None = Field(default=None, max_length=128) + enabled: bool = True + + @model_validator(mode="after") + def validate_scope(self): + if self.scope_type == "global": + self.scope_id = None + elif not (self.scope_id or "").strip(): + raise ValueError("非全局知识必须提供 scope_id") + return self + + +class KnowledgeUpdateRequest(BaseModel): + title: str | None = Field(default=None, min_length=1, max_length=256) + content: str | None = Field(default=None, min_length=1, max_length=50000) + scope_type: ScopeType | None = None + scope_id: str | None = Field(default=None, max_length=128) + enabled: bool | None = None diff --git a/ocai-service/ocai/services/__init__.py b/ocai-service/ocai/services/__init__.py index 98c8e1b..f6eb49e 100644 --- a/ocai-service/ocai/services/__init__.py +++ b/ocai-service/ocai/services/__init__.py @@ -1,9 +1,11 @@ from .chat_service import ChatService +from .knowledge_service import KnowledgeService from .message_converter import MessageConverter from .user_preferences_service import UserPreferencesService __all__ = [ "ChatService", + "KnowledgeService", "MessageConverter", "UserPreferencesService", ] diff --git a/ocai-service/ocai/services/knowledge_service.py b/ocai-service/ocai/services/knowledge_service.py new file mode 100644 index 0000000..84c67f5 --- /dev/null +++ b/ocai-service/ocai/services/knowledge_service.py @@ -0,0 +1,62 @@ +from __future__ import annotations + +import uuid +from typing import Any + +from ocai.repo.interfaces.knowledge import AbstractKnowledgeRepository +from ocai.schemas.knowledge import KnowledgeCreateRequest, KnowledgeUpdateRequest + + +class KnowledgeService: + def __init__(self, repository: AbstractKnowledgeRepository): + self._repository = repository + + def create(self, request: KnowledgeCreateRequest, actor: str) -> dict[str, Any]: + return self._serialize( + self._repository.create( + { + "id": uuid.uuid4().hex, + **request.model_dump(), + "enabled": int(request.enabled), + "created_by": actor, + } + ) + ) + + def list(self, **filters: Any) -> list[dict[str, Any]]: + return [self._serialize(item) for item in self._repository.list_entries(**filters)] + + def update( + self, + entry_id: str, + request: KnowledgeUpdateRequest, + actor: str, + ) -> dict[str, Any] | None: + current = self._repository.get(entry_id) + if current is None or current["created_by"] != actor: + return None + values = request.model_dump(exclude_unset=True) + if "enabled" in values: + values["enabled"] = int(values["enabled"]) + scope_type = values.get("scope_type", current["scope_type"]) + scope_changed = "scope_type" in values and scope_type != current["scope_type"] + scope_id = values.get("scope_id") if scope_changed else values.get("scope_id", current["scope_id"]) + if scope_type == "global": + values["scope_id"] = None + elif not (scope_id or "").strip(): + raise ValueError("非全局知识必须提供 scope_id") + updated = self._repository.update(entry_id, values) + return self._serialize(updated) if updated else None + + def delete(self, entry_id: str, actor: str) -> bool: + current = self._repository.get(entry_id) + return bool(current and current["created_by"] == actor and self._repository.delete(entry_id)) + + @staticmethod + def _serialize(entry: dict[str, Any]) -> dict[str, Any]: + result = dict(entry) + for key in ("created_at", "updated_at"): + if hasattr(result.get(key), "isoformat"): + result[key] = result[key].isoformat() + result["enabled"] = bool(result.get("enabled")) + return result diff --git a/ocai-service/sql/0002_knowledge_entries.down.sql b/ocai-service/sql/0002_knowledge_entries.down.sql new file mode 100644 index 0000000..65dc928 --- /dev/null +++ b/ocai-service/sql/0002_knowledge_entries.down.sql @@ -0,0 +1 @@ +DROP TABLE IF EXISTS knowledge_entries; diff --git a/ocai-service/sql/0002_knowledge_entries.up.sql b/ocai-service/sql/0002_knowledge_entries.up.sql new file mode 100644 index 0000000..7022e59 --- /dev/null +++ b/ocai-service/sql/0002_knowledge_entries.up.sql @@ -0,0 +1,15 @@ +CREATE TABLE IF NOT EXISTS knowledge_entries ( + id VARCHAR(64) NOT NULL, + title VARCHAR(256) NOT NULL, + content TEXT NOT NULL, + scope_type VARCHAR(32) NOT NULL, + scope_id VARCHAR(128) NULL, + enabled INTEGER NOT NULL, + version INTEGER NOT NULL, + created_by VARCHAR(128) NOT NULL, + created_at DATETIME NOT NULL, + updated_at DATETIME NOT NULL, + PRIMARY KEY (id), + INDEX idx_knowledge_scope (scope_type, scope_id, enabled), + INDEX idx_knowledge_creator (created_by, updated_at) +) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci; diff --git a/ocai-service/tests/ai/test_enterprise_knowledge_prompt.py b/ocai-service/tests/ai/test_enterprise_knowledge_prompt.py new file mode 100644 index 0000000..01e6559 --- /dev/null +++ b/ocai-service/tests/ai/test_enterprise_knowledge_prompt.py @@ -0,0 +1,31 @@ +from ocai.ai.prompts.system_builder import SystemPromptBuilder + + +def test_enterprise_knowledge_is_scoped_and_traceable_in_prompt(): + calls = [] + + def provider(**kwargs): + calls.append(kwargs) + return [ + { + "id": "kb-1", + "title": "生产变更规范", + "content": "生产变更必须经过审批。", + } + ] + + builder = SystemPromptBuilder(knowledge_provider=provider) + + prompt = builder.build( + "重启服务", + user_id="alice", + include_tools=False, + include_memory=False, + include_skills=False, + preloaded_context={"project_id": 42, "role": "operator"}, + ) + + assert calls == [{"user_id": "alice", "project_id": "42", "role": "operator"}] + assert "## 企业运维知识" in prompt + assert "生产变更必须经过审批" in prompt + assert "来源: knowledge:kb-1" in prompt diff --git a/ocai-service/tests/api/test_knowledge_endpoints.py b/ocai-service/tests/api/test_knowledge_endpoints.py new file mode 100644 index 0000000..33413ac --- /dev/null +++ b/ocai-service/tests/api/test_knowledge_endpoints.py @@ -0,0 +1,55 @@ +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from ocai.api.v2_route import router +from ocai.permission import UserInfo, get_current_user + + +class _KnowledgeService: + def __init__(self): + self.calls = [] + + def list(self, **filters): + self.calls.append(("list", filters)) + return [] + + def create(self, request, actor): + self.calls.append(("create", actor, request.title)) + return {"id": "k1", "title": request.title, "created_by": actor} + + def update(self, entry_id, request, actor): + self.calls.append(("update", entry_id, actor, request.enabled)) + return {"id": entry_id, "enabled": request.enabled} + + def delete(self, entry_id, actor): + self.calls.append(("delete", entry_id, actor)) + return True + + +def _client(): + app = FastAPI() + service = _KnowledgeService() + app.state.knowledge_service = service + app.include_router(router) + app.dependency_overrides[get_current_user] = lambda: UserInfo(login_name="alice") + return TestClient(app), service + + +def test_knowledge_crud_uses_current_user(): + client, service = _client() + + created = client.post( + "/api/v2/ocai/knowledge", + json={"title": "变更规范", "content": "必须审批"}, + ) + listed = client.get("/api/v2/ocai/knowledge?enabled=true") + updated = client.put("/api/v2/ocai/knowledge/k1", json={"enabled": False}) + deleted = client.delete("/api/v2/ocai/knowledge/k1") + + assert created.status_code == 201 + assert listed.status_code == 200 + assert updated.status_code == 200 + assert deleted.status_code == 200 + assert ("create", "alice", "变更规范") in service.calls + assert ("list", {"scope_type": None, "scope_id": None, "enabled": True, "created_by": "alice"}) in service.calls + assert ("delete", "k1", "alice") in service.calls diff --git a/ocai-service/tests/repo/test_mysql_repos.py b/ocai-service/tests/repo/test_mysql_repos.py index 2f9b6f7..4dea386 100644 --- a/ocai-service/tests/repo/test_mysql_repos.py +++ b/ocai-service/tests/repo/test_mysql_repos.py @@ -20,6 +20,7 @@ from ocai.repo.mysql import ( MysqlAgentRepository, MysqlChatSessionRepository, MySQLEngine, + MysqlKnowledgeRepository, MysqlUserPreferencesRepo, MysqlV2SessionRepository, ) @@ -129,6 +130,51 @@ def test_user_preferences_persist(engine): assert got["enabled_mcp_tools"] == ["x"] +def test_knowledge_entries_crud_and_scope_filtering(engine): + repo = MysqlKnowledgeRepository(engine) + global_entry = repo.create( + { + "id": "global-1", + "title": "生产变更", + "content": "生产变更必须审批", + "scope_type": "global", + "scope_id": None, + "enabled": 1, + "created_by": "alice", + } + ) + repo.create( + { + "id": "project-1", + "title": "项目 SLA", + "content": "四小时内恢复", + "scope_type": "project", + "scope_id": "42", + "enabled": 1, + "created_by": "alice", + } + ) + repo.create( + { + "id": "other-project", + "title": "其他项目", + "content": "不应注入", + "scope_type": "project", + "scope_id": "99", + "enabled": 1, + "created_by": "alice", + } + ) + + applicable = repo.list_applicable(user_id="alice", project_id="42") + + assert [entry["id"] for entry in applicable] == ["global-1", "project-1"] + updated = repo.update("global-1", {"content": "生产变更必须双人审批"}) + assert updated["version"] == global_entry["version"] + 1 + assert repo.delete("global-1") is True + assert repo.get("global-1") is None + + # --------------------------------------------------------------------------- # Agent logs + memory # --------------------------------------------------------------------------- diff --git a/ocai-service/tests/services/test_knowledge_service.py b/ocai-service/tests/services/test_knowledge_service.py new file mode 100644 index 0000000..0a2d909 --- /dev/null +++ b/ocai-service/tests/services/test_knowledge_service.py @@ -0,0 +1,54 @@ +from unittest.mock import MagicMock + +import pytest + +from ocai.schemas.knowledge import KnowledgeCreateRequest, KnowledgeUpdateRequest +from ocai.services.knowledge_service import KnowledgeService + + +def test_creator_owns_knowledge_mutations(): + repo = MagicMock() + repo.create.side_effect = lambda entry: {**entry, "version": 1} + repo.get.return_value = { + "id": "k1", + "created_by": "alice", + "scope_type": "global", + "scope_id": None, + } + repo.update.return_value = { + "id": "k1", + "created_by": "alice", + "scope_type": "global", + "scope_id": None, + "enabled": 1, + "version": 2, + } + service = KnowledgeService(repo) + + created = service.create( + KnowledgeCreateRequest(title="规范", content="必须审批"), + "alice", + ) + updated = service.update("k1", KnowledgeUpdateRequest(enabled=False), "alice") + + assert created["created_by"] == "alice" + assert updated["version"] == 2 + assert service.update("k1", KnowledgeUpdateRequest(enabled=False), "bob") is None + + +def test_scope_change_requires_a_matching_scope_id(): + repo = MagicMock() + repo.get.return_value = { + "id": "k1", + "created_by": "alice", + "scope_type": "project", + "scope_id": "42", + } + service = KnowledgeService(repo) + + with pytest.raises(ValueError, match="scope_id"): + service.update( + "k1", + KnowledgeUpdateRequest(scope_type="user"), + "alice", + ) -- Gitee