diff --git a/LightAgent/__init__.py b/LightAgent/__init__.py index d309612..73a521b 100644 --- a/LightAgent/__init__.py +++ b/LightAgent/__init__.py @@ -68,6 +68,86 @@ ConnectorValidator, validate_connector, ) +from .session import ( + SESSION_SCHEMA_VERSION, + CompactionResult, + ContextBudget, + ContextCompactor, + ContextProjector, + InMemorySessionStore, + JsonlSessionStore, + Session, + SessionCheckpoint, + SessionEvent, + SessionMigrationRegistry, + SessionReplay, + SessionStore, + SqliteSessionStore, + session_migrations, +) +from .capabilities import ( + BaseCapabilityProvider, + BrowserProvider, + CapabilityProvider, + CapabilityRegistry, + CapabilityRisk, + CapabilityScope, + CapabilitySpec, + CredentialProvider, + FileSystemProvider, + InteractionProvider, + LSPProvider, + MemoryProvider, + MemoryProviderAdapter, + ModelProvider, + PermissionSet, + PolicyDecision, + PolicyEngine, + PolicyRequest, + ProviderHealth, + RAGProvider, + RuntimeContext, + SandboxProvider, + ShellProvider, + SubagentProvider, + TelemetryProvider, + TerminalProvider, + ToolProvider, + ToolProviderAdapter, + WebProvider, + WorkflowProvider, +) +from .runtime import ( + AgentInbox, + AgentRuntime, + BudgetExceeded, + BudgetLimits, + BudgetManager, + BudgetUsage, + Goal, + GoalManager, + GoalStatus, + InboxMessage, + InboxMessageStatus, + InboxMessageType, + JobManager, + JobRecord, + JobStatus, + ProgressState, + ProgressTracker, + SubagentManager, + SubagentRecord, +) +from .knowledge import ( + MCPProviderAdapter, + RetrievalDocument, + RetrievalProvider, + RetrievalResult, + SessionSearchProvider, + SkillProviderAdapter, + SqliteFTSRetrievalProvider, + WorkflowProviderAdapter, +) from .builtin_tools.python_executor import ( execute_python_code, execute_python_file, @@ -142,6 +222,78 @@ "ConnectorValidationReport", "ConnectorValidator", "validate_connector", + "SESSION_SCHEMA_VERSION", + "SessionEvent", + "SessionCheckpoint", + "SessionReplay", + "Session", + "SessionStore", + "SessionMigrationRegistry", + "session_migrations", + "InMemorySessionStore", + "JsonlSessionStore", + "SqliteSessionStore", + "ContextProjector", + "ContextBudget", + "ContextCompactor", + "CompactionResult", + "CapabilityScope", + "CapabilityRisk", + "CapabilitySpec", + "ProviderHealth", + "RuntimeContext", + "CapabilityProvider", + "BaseCapabilityProvider", + "ModelProvider", + "ToolProvider", + "FileSystemProvider", + "ShellProvider", + "TerminalProvider", + "BrowserProvider", + "WebProvider", + "LSPProvider", + "MemoryProvider", + "RAGProvider", + "SubagentProvider", + "WorkflowProvider", + "InteractionProvider", + "SandboxProvider", + "CredentialProvider", + "TelemetryProvider", + "PermissionSet", + "PolicyRequest", + "PolicyDecision", + "PolicyEngine", + "CapabilityRegistry", + "ToolProviderAdapter", + "MemoryProviderAdapter", + "InboxMessageType", + "InboxMessageStatus", + "InboxMessage", + "AgentInbox", + "GoalStatus", + "Goal", + "GoalManager", + "BudgetLimits", + "BudgetUsage", + "BudgetExceeded", + "BudgetManager", + "ProgressState", + "ProgressTracker", + "JobStatus", + "JobRecord", + "JobManager", + "SubagentRecord", + "SubagentManager", + "AgentRuntime", + "RetrievalDocument", + "RetrievalResult", + "RetrievalProvider", + "SqliteFTSRetrievalProvider", + "SessionSearchProvider", + "SkillProviderAdapter", + "MCPProviderAdapter", + "WorkflowProviderAdapter", "execute_python_code", "execute_python_file", "execute_python_code_stream", diff --git a/LightAgent/capabilities.py b/LightAgent/capabilities.py new file mode 100644 index 0000000..88578e0 --- /dev/null +++ b/LightAgent/capabilities.py @@ -0,0 +1,685 @@ +"""Capability Provider registry, lifecycle, permissions, and policy decisions.""" + +from __future__ import annotations + +import asyncio +import inspect +import threading +import hashlib +import json +from copy import deepcopy +from dataclasses import asdict, dataclass, field +from enum import Enum +from typing import Any, Awaitable, Callable, Iterable, Protocol + + +class CapabilityScope(str, Enum): + DEFAULT = "default" + RUNTIME = "runtime" + SESSION = "session" + AGENT = "agent" + + +class CapabilityRisk(str, Enum): + READ_ONLY = "L0" + ISOLATED_WRITE = "L1" + SENSITIVE = "L2" + DESTRUCTIVE = "L3" + + +@dataclass(frozen=True) +class CapabilitySpec: + name: str + description: str = "" + risk: CapabilityRisk = CapabilityRisk.READ_ONLY + read: bool = False + write: bool = False + network: bool = False + execute: bool = False + persistent: bool = False + cancellable: bool = False + resumable: bool = False + timeout: float | None = None + output_limit: int | None = None + requires_sandbox: bool = False + requires_approval: bool = False + ui_type: str | None = None + metadata: dict[str, Any] = field(default_factory=dict) + + def __post_init__(self) -> None: + if not self.name: + raise ValueError("CapabilitySpec.name must not be empty") + if self.timeout is not None and self.timeout <= 0: + raise ValueError("CapabilitySpec.timeout must be positive") + if self.output_limit is not None and self.output_limit < 1: + raise ValueError("CapabilitySpec.output_limit must be at least 1") + + def to_dict(self) -> dict[str, Any]: + value = asdict(self) + value["risk"] = self.risk.value + return value + + +@dataclass +class ProviderHealth: + healthy: bool = True + status: str = "ready" + message: str | None = None + degraded_capabilities: list[str] = field(default_factory=list) + metadata: dict[str, Any] = field(default_factory=dict) + + def to_dict(self) -> dict[str, Any]: + return asdict(self) + + +@dataclass +class RuntimeContext: + runtime_id: str | None = None + session_id: str | None = None + agent_id: str | None = None + user_id: str | None = None + turn_id: str | None = None + run_id: str | None = None + permissions: "PermissionSet | None" = None + metadata: dict[str, Any] = field(default_factory=dict) + + +class CapabilityProvider(Protocol): + name: str + version: str + capabilities: dict[str, CapabilitySpec] + + async def mount(self, context: RuntimeContext) -> None: + ... + + async def start(self) -> None: + ... + + async def health(self) -> ProviderHealth: + ... + + async def reload(self, config: dict[str, Any]) -> None: + ... + + async def stop(self) -> None: + ... + + async def unmount(self) -> None: + ... + + +class ModelProvider(CapabilityProvider, Protocol): + pass + + +class ToolProvider(CapabilityProvider, Protocol): + pass + + +class FileSystemProvider(CapabilityProvider, Protocol): + pass + + +class ShellProvider(CapabilityProvider, Protocol): + pass + + +class TerminalProvider(CapabilityProvider, Protocol): + pass + + +class BrowserProvider(CapabilityProvider, Protocol): + pass + + +class WebProvider(CapabilityProvider, Protocol): + pass + + +class LSPProvider(CapabilityProvider, Protocol): + pass + + +class MemoryProvider(CapabilityProvider, Protocol): + pass + + +class RAGProvider(CapabilityProvider, Protocol): + pass + + +class SubagentProvider(CapabilityProvider, Protocol): + pass + + +class WorkflowProvider(CapabilityProvider, Protocol): + pass + + +class InteractionProvider(CapabilityProvider, Protocol): + pass + + +class SandboxProvider(CapabilityProvider, Protocol): + pass + + +class CredentialProvider(CapabilityProvider, Protocol): + pass + + +class TelemetryProvider(CapabilityProvider, Protocol): + pass + + +class BaseCapabilityProvider: + """Small lifecycle implementation for Python-native Providers.""" + + name = "provider" + version = "1" + + def __init__(self, capabilities: Iterable[CapabilitySpec] | None = None): + self.capabilities = {spec.name: spec for spec in capabilities or []} + self.context: RuntimeContext | None = None + self.config: dict[str, Any] = {} + self.mounted = False + self.started = False + + async def mount(self, context: RuntimeContext) -> None: + self.context = context + self.mounted = True + + async def start(self) -> None: + if not self.mounted: + raise RuntimeError(f"provider `{self.name}` must be mounted before start") + self.started = True + + async def health(self) -> ProviderHealth: + return ProviderHealth(healthy=self.started, status="ready" if self.started else "stopped") + + async def reload(self, config: dict[str, Any]) -> None: + self.config = deepcopy(config) + + async def stop(self) -> None: + self.started = False + + async def unmount(self) -> None: + if self.started: + await self.stop() + self.context = None + self.mounted = False + + +@dataclass(frozen=True) +class PermissionSet: + """A capability allowlist that can only be narrowed by descendants.""" + + allowed: frozenset[str] = field(default_factory=frozenset) + denied: frozenset[str] = field(default_factory=frozenset) + max_risk: CapabilityRisk = CapabilityRisk.DESTRUCTIVE + + def allows(self, capability: str, risk: CapabilityRisk = CapabilityRisk.READ_ONLY) -> bool: + if capability in self.denied: + return False + if self.allowed and capability not in self.allowed: + return False + order = { + CapabilityRisk.READ_ONLY: 0, + CapabilityRisk.ISOLATED_WRITE: 1, + CapabilityRisk.SENSITIVE: 2, + CapabilityRisk.DESTRUCTIVE: 3, + } + return order[risk] <= order[self.max_risk] + + def narrow( + self, + *, + allowed: Iterable[str] | None = None, + denied: Iterable[str] | None = None, + max_risk: CapabilityRisk | None = None, + ) -> "PermissionSet": + requested = frozenset(allowed) if allowed is not None else self.allowed + if self.allowed and not requested.issubset(self.allowed): + extra = sorted(requested - self.allowed) + raise ValueError(f"child permissions cannot add capabilities: {extra}") + risk = max_risk or self.max_risk + order = { + CapabilityRisk.READ_ONLY: 0, + CapabilityRisk.ISOLATED_WRITE: 1, + CapabilityRisk.SENSITIVE: 2, + CapabilityRisk.DESTRUCTIVE: 3, + } + if order[risk] > order[self.max_risk]: + raise ValueError("child permissions cannot increase max_risk") + return PermissionSet( + allowed=requested, + denied=self.denied | frozenset(denied or []), + max_risk=risk, + ) + + +@dataclass +class PolicyRequest: + capability: CapabilitySpec + provider_name: str + arguments: dict[str, Any] = field(default_factory=dict) + context: RuntimeContext = field(default_factory=RuntimeContext) + + +@dataclass +class PolicyDecision: + allowed: bool + reason: str | None = None + requires_approval: bool = False + arguments: dict[str, Any] | None = None + metadata: dict[str, Any] = field(default_factory=dict) + + @classmethod + def allow(cls, arguments: dict[str, Any] | None = None, **metadata: Any) -> "PolicyDecision": + return cls(allowed=True, arguments=arguments, metadata=metadata) + + @classmethod + def block(cls, reason: str, **metadata: Any) -> "PolicyDecision": + return cls(allowed=False, reason=reason, metadata=metadata) + + @classmethod + def approval(cls, reason: str | None = None, **metadata: Any) -> "PolicyDecision": + return cls(allowed=False, reason=reason, requires_approval=True, metadata=metadata) + + +PolicyCallable = Callable[[PolicyRequest], PolicyDecision | bool | dict[str, Any] | None | Awaitable[Any]] + + +class PolicyEngine: + """Ordered fail-closed policy evaluation for capability execution.""" + + def __init__(self, policies: Iterable[PolicyCallable] | None = None, *, fail_closed: bool = True): + self.policies = list(policies or []) + self.fail_closed = fail_closed + + def add(self, policy: PolicyCallable) -> None: + self.policies.append(policy) + + async def evaluate(self, request: PolicyRequest) -> PolicyDecision: + permissions = request.context.permissions + if permissions and not permissions.allows(request.capability.name, request.capability.risk): + return PolicyDecision.block(f"capability `{request.capability.name}` is outside the permission snapshot") + if request.capability.requires_approval: + return PolicyDecision.approval(f"capability `{request.capability.name}` requires approval") + + arguments = deepcopy(request.arguments) + collected_metadata: dict[str, Any] = {} + for policy in self.policies: + try: + result = policy(PolicyRequest( + capability=request.capability, + provider_name=request.provider_name, + arguments=deepcopy(arguments), + context=request.context, + )) + if inspect.isawaitable(result): + result = await result + decision = self._coerce(result, arguments) + except Exception as error: + if self.fail_closed: + return PolicyDecision.block( + f"policy `{getattr(policy, '__name__', policy.__class__.__name__)}` failed: {type(error).__name__}" + ) + collected_metadata.setdefault("policy_errors", []).append(type(error).__name__) + continue + collected_metadata.update(decision.metadata) + if not decision.allowed: + decision.metadata = {**collected_metadata, **decision.metadata} + return decision + if decision.arguments is not None: + arguments = deepcopy(decision.arguments) + return PolicyDecision.allow(arguments, **collected_metadata) + + def evaluate_sync(self, request: PolicyRequest) -> PolicyDecision: + return _run_sync(self.evaluate(request)) + + @staticmethod + def _coerce(value: Any, arguments: dict[str, Any]) -> PolicyDecision: + if value is None or value is True: + return PolicyDecision.allow(arguments) + if value is False: + return PolicyDecision.block("policy denied the capability") + if isinstance(value, PolicyDecision): + return value + if isinstance(value, dict): + if value.get("requires_approval"): + return PolicyDecision.approval(value.get("reason"), **dict(value.get("metadata") or {})) + allowed = bool(value.get("allowed", True)) + return PolicyDecision( + allowed=allowed, + reason=value.get("reason"), + arguments=value.get("arguments", arguments), + metadata=dict(value.get("metadata") or {}), + ) + raise TypeError("policy must return PolicyDecision, bool, dict, or None") + + +@dataclass +class ProviderRegistration: + provider: CapabilityProvider + scope: CapabilityScope + owner_id: str | None = None + order: int = 0 + + +class CapabilityRegistry: + """Resolve Providers by capability and runtime/session/agent scope.""" + + _precedence = { + CapabilityScope.DEFAULT: 0, + CapabilityScope.RUNTIME: 1, + CapabilityScope.SESSION: 2, + CapabilityScope.AGENT: 3, + } + + def __init__( + self, + *, + policy_engine: PolicyEngine | None = None, + audit: Callable[[str, dict[str, Any]], Any] | None = None, + ): + self.policy_engine = policy_engine or PolicyEngine() + self.audit = audit + self._registrations: list[ProviderRegistration] = [] + self._conflicts: list[dict[str, Any]] = [] + self._counter = 0 + + def register( + self, + provider: CapabilityProvider, + *, + scope: CapabilityScope | str = CapabilityScope.RUNTIME, + owner_id: str | None = None, + ) -> CapabilityProvider: + resolved_scope = CapabilityScope(scope) + if resolved_scope in {CapabilityScope.SESSION, CapabilityScope.AGENT} and not owner_id: + raise ValueError(f"owner_id is required for {resolved_scope.value} scope") + if any( + item.provider.name == provider.name and item.scope == resolved_scope and item.owner_id == owner_id + for item in self._registrations + ): + raise ValueError( + f"provider `{provider.name}` is already registered for {resolved_scope.value}:{owner_id or '*'}" + ) + self._counter += 1 + for existing in self._registrations: + overlap = sorted(set(existing.provider.capabilities) & set(provider.capabilities)) + if overlap and existing.scope == resolved_scope and existing.owner_id == owner_id: + self._conflicts.append({ + "capabilities": overlap, + "winner": provider.name, + "shadowed": existing.provider.name, + "scope": resolved_scope.value, + "owner_id": owner_id, + "rule": "latest registration wins within an equal scope", + }) + self._registrations.append(ProviderRegistration(provider, resolved_scope, owner_id, self._counter)) + return provider + + def conflicts(self) -> list[dict[str, Any]]: + return deepcopy(self._conflicts) + + async def mount(self, context: RuntimeContext) -> None: + for item in self._matching(context): + await item.provider.mount(context) + await item.provider.start() + self._audit("provider.started", item, context=context) + + async def unregister( + self, + name: str, + *, + scope: CapabilityScope | str | None = None, + owner_id: str | None = None, + ) -> bool: + resolved_scope = CapabilityScope(scope) if scope is not None else None + matches = [ + item for item in self._registrations + if item.provider.name == name + and (resolved_scope is None or item.scope == resolved_scope) + and (owner_id is None or item.owner_id == owner_id) + ] + for item in reversed(matches): + await item.provider.stop() + await item.provider.unmount() + self._registrations.remove(item) + self._audit("provider.unregistered", item) + return bool(matches) + + def resolve(self, capability: str, context: RuntimeContext | None = None) -> CapabilityProvider: + runtime_context = context or RuntimeContext() + matches = [ + item for item in self._matching(runtime_context) + if capability in item.provider.capabilities + ] + if not matches: + raise LookupError(f"no Provider registered for capability `{capability}`") + matches.sort(key=lambda item: (self._precedence[item.scope], item.order), reverse=True) + return matches[0].provider + + def get(self, name: str, context: RuntimeContext | None = None) -> CapabilityProvider: + matches = [item for item in self._matching(context or RuntimeContext()) if item.provider.name == name] + if not matches: + raise LookupError(f"provider `{name}` is not registered") + matches.sort(key=lambda item: (self._precedence[item.scope], item.order), reverse=True) + return matches[0].provider + + def list(self, context: RuntimeContext | None = None) -> list[dict[str, Any]]: + values = [] + for item in self._matching(context or RuntimeContext()): + values.append({ + "name": item.provider.name, + "version": item.provider.version, + "scope": item.scope.value, + "owner_id": item.owner_id, + "capabilities": [spec.to_dict() for spec in item.provider.capabilities.values()], + }) + return values + + async def health(self, context: RuntimeContext | None = None) -> dict[str, ProviderHealth]: + result = {} + for item in self._matching(context or RuntimeContext()): + result[item.provider.name] = await item.provider.health() + return result + + async def reload(self, name: str, config: dict[str, Any], context: RuntimeContext | None = None) -> None: + provider = self.get(name, context) + await provider.reload(config) + + async def stop(self, context: RuntimeContext | None = None) -> None: + for item in reversed(self._matching(context or RuntimeContext())): + await item.provider.stop() + await item.provider.unmount() + self._audit("provider.stopped", item, context=context) + + async def invoke( + self, + capability: str, + arguments: dict[str, Any] | None = None, + *, + context: RuntimeContext | None = None, + ) -> Any: + runtime_context = context or RuntimeContext() + provider = self.resolve(capability, runtime_context) + spec = provider.capabilities[capability] + decision = await self.policy_engine.evaluate(PolicyRequest( + capability=spec, + provider_name=provider.name, + arguments=arguments or {}, + context=runtime_context, + )) + self._audit("policy.decision", self._registration_for(provider), context=runtime_context, extra={ + "capability": capability, + "allowed": decision.allowed, + "requires_approval": decision.requires_approval, + "reason": decision.reason, + }) + if not decision.allowed: + state = "requires approval" if decision.requires_approval else "was denied" + raise PermissionError(f"capability `{capability}` {state}: {decision.reason or 'policy decision'}") + invoke = getattr(provider, "invoke", None) + if not callable(invoke): + raise TypeError(f"provider `{provider.name}` does not implement invoke()") + result = invoke(capability, **(decision.arguments or arguments or {})) + if inspect.isawaitable(result): + result = await asyncio.wait_for(result, timeout=spec.timeout) if spec.timeout else await result + if spec.output_limit is not None and len(str(result)) > spec.output_limit: + result = str(result)[:spec.output_limit] + return result + + def _matching(self, context: RuntimeContext) -> list[ProviderRegistration]: + return [ + item for item in self._registrations + if item.scope in {CapabilityScope.DEFAULT, CapabilityScope.RUNTIME} + or (item.scope == CapabilityScope.SESSION and item.owner_id == context.session_id) + or (item.scope == CapabilityScope.AGENT and item.owner_id == context.agent_id) + ] + + def _registration_for(self, provider: CapabilityProvider) -> ProviderRegistration: + return next(item for item in self._registrations if item.provider is provider) + + def _audit( + self, + event_type: str, + item: ProviderRegistration, + *, + context: RuntimeContext | None = None, + extra: dict[str, Any] | None = None, + ) -> None: + if not self.audit: + return + self.audit(event_type, { + "provider": item.provider.name, + "provider_version": item.provider.version, + "scope": item.scope.value, + "owner_id": item.owner_id, + "session_id": context.session_id if context else None, + "agent_id": context.agent_id if context else None, + "configuration_digest": self._configuration_digest(item.provider), + **(extra or {}), + }) + + @staticmethod + def _configuration_digest(provider: CapabilityProvider) -> str: + config = getattr(provider, "config", {}) + safe = { + key: "[redacted]" if any(token in str(key).lower() for token in ("key", "token", "secret", "password")) else value + for key, value in dict(config or {}).items() + } + rendered = json.dumps(safe, ensure_ascii=True, sort_keys=True, default=repr) + return hashlib.sha256(rendered.encode("utf-8")).hexdigest() + + +class ToolProviderAdapter(BaseCapabilityProvider): + name = "tools" + + def __init__(self, tool_registry: Any): + self.tool_registry = tool_registry + specs = [ + CapabilitySpec( + name=f"tool.{name}", + description=str(info.get("tool_description", "")), + execute=True, + risk=CapabilityRisk.SENSITIVE, + cancellable=True, + ) + for name, info in tool_registry.function_info.items() + ] + super().__init__(specs) + + async def invoke(self, capability: str, **arguments: Any) -> Any: + name = capability.removeprefix("tool.") + from .tools import AsyncToolDispatcher + + dispatcher = AsyncToolDispatcher(self.tool_registry.function_mappings, self.tool_registry.function_info) + return await dispatcher.dispatch(name, arguments) + + +class MemoryProviderAdapter(BaseCapabilityProvider): + name = "memory" + + def __init__(self, backend: Any): + self.backend = backend + super().__init__([ + CapabilitySpec("memory.retrieve", read=True), + CapabilitySpec( + "memory.store", + write=True, + persistent=True, + risk=CapabilityRisk.ISOLATED_WRITE, + ), + ]) + + async def invoke(self, capability: str, **arguments: Any) -> Any: + if capability == "memory.retrieve": + result = self.backend.retrieve(arguments["query"], arguments["user_id"]) + elif capability == "memory.store": + result = self.backend.store(arguments["data"], arguments["user_id"]) + else: + raise LookupError(capability) + if inspect.isawaitable(result): + return await result + return result + + +def _run_sync(awaitable: Awaitable[Any]) -> Any: + try: + asyncio.get_running_loop() + except RuntimeError: + return asyncio.run(awaitable) + + result: list[Any] = [] + error: list[BaseException] = [] + + def runner() -> None: + try: + result.append(asyncio.run(awaitable)) + except BaseException as exc: # pragma: no cover - re-raised below + error.append(exc) + + thread = threading.Thread(target=runner, daemon=True) + thread.start() + thread.join() + if error: + raise error[0] + return result[0] + + +__all__ = [ + "CapabilityScope", + "CapabilityRisk", + "CapabilitySpec", + "ProviderHealth", + "RuntimeContext", + "CapabilityProvider", + "ModelProvider", + "ToolProvider", + "FileSystemProvider", + "ShellProvider", + "TerminalProvider", + "BrowserProvider", + "WebProvider", + "LSPProvider", + "MemoryProvider", + "RAGProvider", + "SubagentProvider", + "WorkflowProvider", + "InteractionProvider", + "SandboxProvider", + "CredentialProvider", + "TelemetryProvider", + "BaseCapabilityProvider", + "PermissionSet", + "PolicyRequest", + "PolicyDecision", + "PolicyEngine", + "ProviderRegistration", + "CapabilityRegistry", + "ToolProviderAdapter", + "MemoryProviderAdapter", +] diff --git a/LightAgent/core.py b/LightAgent/core.py index 219e0fc..592c582 100644 --- a/LightAgent/core.py +++ b/LightAgent/core.py @@ -43,6 +43,24 @@ from .mcp_client_manager import MCPClientManager from .skills import SkillManager from .skill_tools import create_skill_tools +from .capabilities import ( + CapabilityRegistry, + CapabilityRisk, + CapabilityScope, + CapabilitySpec, + MemoryProviderAdapter, + PolicyEngine, + PolicyRequest, + ToolProviderAdapter, +) +from .runtime import AgentRuntime, BudgetExceeded, BudgetLimits +from .session import ( + ContextBudget, + ContextCompactor, + ContextProjector, + Session, + SessionStore, +) # 新增:导入内置工具 from .builtin_tools.python_executor import ( execute_python_code, @@ -119,6 +137,12 @@ def __init__( tool_guardrails: List[Callable[..., Any]] | None = None, # 工具调用安全策略 output_guardrails: List[Callable[..., Any]] | None = None, # 输出安全策略 hooks: List[Callable[..., Any] | PolicyHook] | None = None, # 运行期 hook / middleware + session_store: SessionStore | None = None, # v0.10 Session 持久化后端 + capability_registry: CapabilityRegistry | None = None, # v0.10 能力注册表 + policy_engine: PolicyEngine | None = None, # v0.10 统一能力策略 + budget_limits: BudgetLimits | None = None, # v0.10 长任务预算 + context_budget: ContextBudget | None = None, # 可选模型上下文预算 + context_compactor: ContextCompactor | None = None, # 可选上下文压缩器 debug: bool = False, # 是否启用调试模式 log_level: str = "INFO", # 日志级别(INFO, DEBUG, ERROR) log_file: Optional[str] = None, # 日志文件路径 @@ -148,6 +172,12 @@ def __init__( :param tool_guardrails: 工具调用安全策略列表,返回 False、原因字符串、dict 或 GuardrailDecision 可阻止工具执行。 :param output_guardrails: 输出安全策略列表,返回 False、原因字符串、dict 或 GuardrailDecision 可阻止非流式输出。 :param hooks: 运行期 hook 列表,可观察、替换或阻断指定生命周期阶段。 + :param session_store: 可选 SessionStore;默认使用进程内存储且不影响旧调用方式。 + :param capability_registry: 可选 CapabilityRegistry,用于统一 Provider 生命周期与策略。 + :param policy_engine: 可选能力策略引擎;传入 registry 时应优先配置 registry 自身策略。 + :param budget_limits: 可选模型调用、工具调用、token、时间和成本预算。 + :param context_budget: 可选上下文 token 预算。 + :param context_compactor: 可选两阶段上下文压缩器。 :param debug: 是否启用调试模式。 :param log_level: 日志级别(INFO, DEBUG, ERROR)。 :param log_file: 日志文件路径。 @@ -199,6 +229,8 @@ def __init__( self._parent_trace_id: str | None = None self._run_group_id: str | None = None self._memory_promotion_candidates: list[MemoryCandidate] = [] + self.context_budget = context_budget + self.context_compactor = context_compactor or ContextCompactor() # 确保 log 目录存在 log_dir = 'logs' if not os.path.exists(log_dir): @@ -283,6 +315,21 @@ def __init__( self.tracetools = tracetools self.chat_params = {} # history 存储器 self._trace_recorder = TraceRecorder(enabled=False) + registry = capability_registry or CapabilityRegistry(policy_engine=policy_engine) + if capability_registry is not None and policy_engine is not None: + registry.policy_engine = policy_engine + self.runtime = AgentRuntime( + session_store=session_store, + capability_registry=registry, + budget_limits=budget_limits, + ) + self.capability_registry = self.runtime.registry + self.capability_registry.audit = self._record_session_event + self._register_builtin_providers() + self._current_session_id: str | None = None + self._current_turn_id: str | None = None + self._session_turn_finished = True + self._session_store_error: str | None = None def _initialize_clients(self, tracetools, tot_api_key, tot_base_url, tot_model): """初始化 OpenAI 客户端""" @@ -436,6 +483,7 @@ def run( parent_trace_id: str | None = None, run_group_id: str | None = None, approval_id: str | None = None, + session_id: str | None = None, ) -> Union[Generator[str, None, None], str, RunResult]: """ 运行代理,处理用户输入。 @@ -455,6 +503,7 @@ def run( :param parent_trace_id: 可选父 trace,用于 LightFlow、LightSwarm 或应用层嵌套调用。 :param run_group_id: 可选运行组 ID,用于把多个 sibling traces 归到同一任务。 :param approval_id: 可选人工审批请求 ID,仅传给运行期 hooks,不发送给模型服务。 + :param session_id: 可选持久化 Session ID。省略时每次调用创建独立的内存 Session。 :return: 代理的回复。 """ if result_format not in ("str", "object", "dict", "event"): @@ -499,6 +548,14 @@ def run( self._memory_write_count = 0 self._memory_write_fingerprints = set() self._memory_promotion_candidates = [] + self._budget_error = None + projected_history = self._begin_session_run( + session_id=session_id, + query=query, + user_id=str(user_id), + metadata=dict(metadata or {}), + stream=stream, + ) if self.debug and hasattr(self, 'logger'): # 仅在 debug=True 且 logger 存在时记录日志 self.logger.set_traceid(traceid) self.log("INFO", "run_start", {"query": query, "user_id": user_id, "stream": stream}) @@ -554,7 +611,7 @@ def run( query = input_decision.value # 初始化历史记录 - history = history or [] + history = projected_history if history is None and session_id is not None else (history or []) # 处理运行时传入的工具 runtime_tools = [] @@ -675,11 +732,196 @@ def run( return result return self._format_run_result(result, result_format, traceid) + async def arun(self, query: str, **kwargs: Any) -> Any: + """Run without blocking the caller's event loop. + + Sync model and tool clients are isolated in a worker thread for v0.9 + compatibility. Native async Providers and Jobs remain on the caller's + event loop through ``CapabilityRegistry`` and ``AgentRuntime``. + """ + stream = bool(kwargs.get("stream", False)) + if not stream: + return await asyncio.to_thread(self.run, query, **kwargs) + return self.astream(query, **kwargs) + + async def astream(self, query: str, **kwargs: Any) -> AsyncGenerator[Any, None]: + """Expose the legacy streaming generator as a cancellable async iterator.""" + options = dict(kwargs) + options["stream"] = True + stream_result = await asyncio.to_thread(self.run, query, **options) + sentinel = object() + try: + while True: + chunk = await asyncio.to_thread(next, stream_result, sentinel) + if chunk is sentinel: + return + yield chunk + finally: + close = getattr(stream_result, "close", None) + if callable(close): + await asyncio.to_thread(close) + + def get_session(self, session_id: str | None = None) -> Session | None: + """Return a detached copy of the current or requested Session.""" + resolved = session_id or self._current_session_id + return self.runtime.session_store.get(resolved) if resolved else None + + def export_session(self, session_id: str | None = None) -> Dict[str, Any]: + session = self.get_session(session_id) + if session is None: + raise KeyError(session_id or "current session") + return session.to_dict() + + def replay_session(self, session_id: str | None = None) -> Dict[str, Any]: + session = self.get_session(session_id) + if session is None: + raise KeyError(session_id or "current session") + return session.replay().to_dict() + + def checkpoint_session( + self, + label: str | None = None, + metadata: Dict[str, Any] | None = None, + ) -> Dict[str, Any]: + return self.runtime.checkpoint(label, metadata).to_dict() + + def fork_session( + self, + *, + through_sequence: int | None = None, + metadata: Dict[str, Any] | None = None, + ) -> Session: + return self.runtime.fork(through_sequence=through_sequence, metadata=metadata) + + def compact_session( + self, + *, + max_messages: int = 20, + session_id: str | None = None, + ) -> Dict[str, Any]: + session = self.get_session(session_id) + if session is None: + raise KeyError(session_id or "current session") + result = self.context_compactor.compact(ContextProjector().messages(session), max_messages=max_messages) + target = self.runtime.session if self.runtime.session and self.runtime.session.session_id == session.session_id else session + target.append("context.compacted", { + "messages": result.messages, + "removed_count": result.removed_count, + "summary": result.summary, + "spilled": result.spilled, + }) + self.runtime.session_store.save(target) + return { + "messages": result.messages, + "removed_count": result.removed_count, + "summary": result.summary, + "spilled": result.spilled, + } + + def pause_session(self, reason: str | None = None) -> None: + self._record_session_event("session.paused", {"reason": reason}) + + def resume_session(self, reason: str | None = None) -> None: + self._record_session_event("session.resumed", {"reason": reason}) + + def cancel_session(self, reason: str | None = None) -> None: + self._record_session_event("session.cancelled", {"reason": reason}) + + async def invoke_capability(self, capability: str, **arguments: Any) -> Any: + """Invoke one registered capability through Policy and audit.""" + return await self.capability_registry.invoke( + capability, + arguments, + context=self.runtime.context, + ) + + def _begin_session_run( + self, + *, + session_id: str | None, + query: str, + user_id: str, + metadata: Dict[str, Any], + stream: bool, + ) -> List[Dict[str, Any]]: + self.runtime.open_session( + session_id, + metadata=metadata, + agent_id=self.name, + user_id=user_id, + ) + current = self.runtime.session + projected = ContextProjector().messages(current) if current is not None else [] + self._current_session_id = current.session_id if current is not None else None + self._current_turn_id = uuid4().hex + self._session_turn_finished = False + self._session_store_error = None + self.runtime.context.turn_id = self._current_turn_id + self.runtime.context.run_id = self._current_run_id + self.runtime.context.metadata = deepcopy(metadata) + self._record_session_event("turn.started", { + "query": query, + "stream": stream, + "trace_id": self.traceid, + }) + self._record_session_event("message.received", {"role": "user", "content": query}) + return projected + + def _register_builtin_providers(self) -> None: + existing = {item["name"] for item in self.capability_registry.list()} + if "tools" not in existing: + self.capability_registry.register(ToolProviderAdapter(self.tool_registry), scope=CapabilityScope.RUNTIME) + if self.memory is not None and "memory" not in existing: + self.capability_registry.register(MemoryProviderAdapter(self.memory), scope=CapabilityScope.RUNTIME) + + def _record_session_event(self, event_type: str, data: Dict[str, Any] | None = None) -> Any: + if not getattr(self, "runtime", None) or self.runtime.session is None: + return None + try: + session = self.runtime.session + event = session.append( + event_type, + data or {}, + turn_id=getattr(self, "_current_turn_id", None), + run_id=getattr(self, "_current_run_id", None), + agent_id=self.name, + ) + self.runtime.session_store.save(session) + return event + except Exception as error: + self._session_store_error = f"{type(error).__name__}: {error}" + self.log("ERROR", "session_store_error", {"error": self._session_store_error}) + return None + def _record_trace(self, event_type: str, data: Dict[str, Any] | None = None): - """Record a trace event when tracing is enabled.""" + """Record compatible Trace and durable Session views from one event.""" + trace_data = deepcopy(data or {}) + session_type = { + "run_start": "run.started", + "model_request": "model.requested", + "model_response": "assistant.completed", + "tool_call": "tool.requested", + "tool_result": "tool.completed", + "guardrail_block": "policy.blocked", + "hook_block": "policy.blocked", + "error": "error.recorded", + "handoff": "handoff.requested", + "run_end": "run.completed" if trace_data.get("success") else "run.failed", + }.get(event_type, "trace.recorded") + session_data = { + **trace_data, + "trace_type": event_type, + "trace_data": trace_data, + "trace_id": getattr(self, "traceid", None), + "parent_trace_id": getattr(self, "_parent_trace_id", None), + "run_group_id": getattr(self, "_run_group_id", None), + } + if event_type == "model_request": + session_data["messages"] = deepcopy(getattr(self, "chat_params", {}).get("messages", [])) + self._record_session_event(session_type, session_data) recorder = getattr(self, "_trace_recorder", None) if recorder: - return recorder.record(event_type, data) + return recorder.record(event_type, trace_data) return None def _run_hooks( @@ -746,6 +988,20 @@ def _finish_run( extra: Dict[str, Any] | None = None, ) -> HookDecision: """Close a run with on_error/after_run hooks and a run_end trace event.""" + started_at = getattr(self, "_run_started_at", None) + duration_seconds = (time.perf_counter() - started_at) if started_at is not None else 0.0 + budget_error = getattr(self, "_budget_error", None) + if budget_error: + success = False + error = error or budget_error + stage = stage or "budget" + try: + if duration_seconds: + self.runtime.budget.consume(seconds=duration_seconds) + except BudgetExceeded as exceeded: + success = False + error = error or format_error_code("LA-BUDGET", str(exceeded)) + stage = stage or "budget" payload = { "success": success, "content": content, @@ -769,7 +1025,6 @@ def _finish_run( run_end["stage"] = stage if extra: run_end.update(extra) - started_at = getattr(self, "_run_started_at", None) if started_at is not None: run_end["duration_ms"] = round((time.perf_counter() - started_at) * 1000, 3) run_end["model_request_count"] = int(getattr(self, "_model_request_count", 0)) @@ -779,6 +1034,12 @@ def _finish_run( if usage: run_end["usage"] = usage self._record_trace("run_end", run_end) + if not getattr(self, "_session_turn_finished", True): + self._record_session_event("turn.completed" if success else "turn.failed", { + **run_end, + "content": content, + }) + self._session_turn_finished = True return decision def _notify_error(self, *, stage: str, error: str, **payload: Any) -> HookDecision: @@ -795,6 +1056,7 @@ def _format_hook_error(phase: str, reason: str | None = None) -> str: return format_error_code("LA-HOOK", details) def _prepare_model_request(self) -> str | None: + self._compact_model_context() decision = self._run_hooks("before_model_request", {"params": deepcopy(self.chat_params)}) if decision.action == HOOK_BLOCK: return self._format_hook_error("before_model_request", decision.reason) @@ -802,6 +1064,10 @@ def _prepare_model_request(self) -> str | None: params = decision.payload.get("params", decision.payload) if isinstance(params, dict): self.chat_params = params + try: + self.runtime.budget.consume(model_calls=1) + except BudgetExceeded as error: + return format_error_code("LA-BUDGET", str(error)) self._model_request_count = int(getattr(self, "_model_request_count", 0)) + 1 request_data = self._build_model_request_trace(self.chat_params) request_data["request_index"] = self._model_request_count @@ -809,6 +1075,19 @@ def _prepare_model_request(self) -> str | None: self._pending_model_started_at = time.perf_counter() return None + def _compact_model_context(self) -> None: + messages = self.chat_params.get("messages", []) + if not self.context_budget or self.context_budget.fits(messages): + return + compacted = self.context_compactor.compact_to_budget(messages, self.context_budget) + self.chat_params["messages"] = compacted.messages + self._record_session_event("context.compacted", { + "removed_count": compacted.removed_count, + "summary": compacted.summary, + "spilled": compacted.spilled, + "estimated_tokens": self.context_budget.estimate(compacted.messages), + }) + def _complete_model_request( self, *, @@ -833,6 +1112,17 @@ def _complete_model_request( event.data["usage"] = usage if error is not None: event.data["error_type"] = error.__class__.__name__ + self._record_session_event("model.failed" if error is not None else "model.completed", { + "request_index": int(getattr(self, "_model_request_count", 0)), + "latency_ms": latency_ms, + "usage": usage, + "error_type": error.__class__.__name__ if error is not None else None, + }) + if usage: + try: + self.runtime.budget.consume(tokens=int(usage.get("total_tokens", 0))) + except BudgetExceeded as exceeded: + self._budget_error = format_error_code("LA-BUDGET", str(exceeded)) self._last_model_trace_event = event self._pending_model_trace_event = None self._pending_model_started_at = None @@ -934,6 +1224,38 @@ def _apply_tool_guardrails(self, tool_name: str, arguments: Dict[str, Any]) -> t return arguments, error_msg def _prepare_tool_call(self, tool_name: str, arguments: Dict[str, Any]) -> tuple[Dict[str, Any], str | None]: + spec = CapabilitySpec( + name=f"tool.{tool_name}", + description=str(self.tool_registry.function_info.get(tool_name, {}).get("tool_description", "")), + execute=True, + risk=CapabilityRisk.SENSITIVE, + cancellable=True, + ) + policy_decision = self.capability_registry.policy_engine.evaluate_sync(PolicyRequest( + capability=spec, + provider_name="tools", + arguments=deepcopy(arguments), + context=self.runtime.context, + )) + self._record_session_event("policy.decision", { + "provider": "tools", + "capability": spec.name, + "allowed": policy_decision.allowed, + "requires_approval": policy_decision.requires_approval, + "reason": policy_decision.reason, + }) + if not policy_decision.allowed: + state = "requires approval" if policy_decision.requires_approval else "was denied" + return arguments, format_error_code( + "LA-POLICY", + f"capability `{spec.name}` {state}: {policy_decision.reason or 'policy decision'}", + ) + if policy_decision.arguments is not None: + arguments = policy_decision.arguments + try: + self.runtime.budget.consume(tool_calls=1) + except BudgetExceeded as error: + return arguments, format_error_code("LA-BUDGET", str(error)) arguments, guardrail_error = self._apply_tool_guardrails(tool_name, arguments) if guardrail_error: return arguments, guardrail_error diff --git a/LightAgent/errors.py b/LightAgent/errors.py index 4e720f8..df654c2 100644 --- a/LightAgent/errors.py +++ b/LightAgent/errors.py @@ -108,6 +108,36 @@ def __str__(self) -> str: "A configured runtime hook blocked this operation.", "Review the hook decision, configured phase, and application policy before retrying.", ), + "LA-POLICY": LightAgentErrorInfo( + "LA-POLICY", + "The unified capability policy denied this operation.", + "Review Provider scope, permission snapshots, approval requirements, and Policy decisions.", + ), + "LA-BUDGET": LightAgentErrorInfo( + "LA-BUDGET", + "The configured runtime budget was exhausted.", + "Increase the relevant budget or resume with a narrower goal, tool set, or context.", + ), + "LA-SESSION": LightAgentErrorInfo( + "LA-SESSION", + "Session persistence or replay failed.", + "Inspect the SessionStore, event sequence, schema version, and latest checkpoint.", + ), + "LA-PROVIDER": LightAgentErrorInfo( + "LA-PROVIDER", + "A capability Provider failed or is unavailable.", + "Inspect Provider health, lifecycle state, configuration digest, and optional dependencies.", + ), + "LA-APPROVAL": LightAgentErrorInfo( + "LA-APPROVAL", + "This operation is waiting for explicit approval.", + "Resolve the approval request and resume the same Session or workflow run.", + ), + "LA-CANCELLED": LightAgentErrorInfo( + "LA-CANCELLED", + "The run, Job, workflow, or Session was cancelled.", + "Inspect the cancellation event and resume from a checkpoint when the operation is resumable.", + ), "LA-UNKNOWN": LightAgentErrorInfo( "LA-UNKNOWN", "An unexpected LightAgent error occurred.", diff --git a/LightAgent/flow.py b/LightAgent/flow.py index fe2c58a..cfefe6b 100644 --- a/LightAgent/flow.py +++ b/LightAgent/flow.py @@ -8,6 +8,7 @@ from __future__ import annotations import json +import asyncio import time from concurrent.futures import ThreadPoolExecutor, TimeoutError from dataclasses import asdict, dataclass, field @@ -258,6 +259,10 @@ def run( run_group_id=run_group_id, ) + async def arun(self, query: str, **kwargs: Any) -> LightFlowResult | str | dict[str, Any]: + """Run a workflow without blocking the caller's event loop.""" + return await asyncio.to_thread(self.run, query, **kwargs) + def resume( self, run_id: str, diff --git a/LightAgent/knowledge.py b/LightAgent/knowledge.py new file mode 100644 index 0000000..a526757 --- /dev/null +++ b/LightAgent/knowledge.py @@ -0,0 +1,443 @@ +"""Optional knowledge Providers for RAG, Session search, Skills, MCP, and LightFlow.""" + +from __future__ import annotations + +import inspect +import json +import sqlite3 +from copy import deepcopy +from dataclasses import asdict, dataclass, field +from pathlib import Path +from typing import Any, Iterable, Protocol +from uuid import uuid4 + +from .capabilities import ( + BaseCapabilityProvider, + CapabilityRisk, + CapabilitySpec, + ProviderHealth, +) +from .session import SessionStore, _utc_now + + +@dataclass +class RetrievalDocument: + content: str + title: str | None = None + source: str | None = None + document_id: str = field(default_factory=lambda: uuid4().hex) + scope: str = "workspace" + owner_id: str | None = None + tenant_id: str | None = None + created_at: str = field(default_factory=_utc_now) + metadata: dict[str, Any] = field(default_factory=dict) + + def to_dict(self) -> dict[str, Any]: + return asdict(self) + + +@dataclass +class RetrievalResult: + document_id: str + chunk_id: str + content: str + citation_id: str + title: str | None = None + source: str | None = None + position: int = 0 + score: float | None = None + scope: str = "workspace" + owner_id: str | None = None + tenant_id: str | None = None + metadata: dict[str, Any] = field(default_factory=dict) + + def to_dict(self) -> dict[str, Any]: + return asdict(self) + + +class RetrievalProvider(Protocol): + def ingest(self, document: RetrievalDocument) -> str: + ... + + def remove(self, document_id: str) -> bool: + ... + + def list_documents(self) -> list[RetrievalDocument]: + ... + + def search(self, query: str, *, limit: int = 5, **filters: Any) -> list[RetrievalResult]: + ... + + def read_chunk(self, chunk_id: str) -> RetrievalResult | None: + ... + + def reindex(self) -> None: + ... + + +class SqliteFTSRetrievalProvider(BaseCapabilityProvider): + """Dependency-free SQLite FTS5 retrieval with source-aware chunks.""" + + name = "sqlite-fts5-rag" + version = "1" + + def __init__(self, path: str | Path, *, chunk_size: int = 1200, chunk_overlap: int = 120): + if chunk_size < 100: + raise ValueError("chunk_size must be at least 100") + if chunk_overlap < 0 or chunk_overlap >= chunk_size: + raise ValueError("chunk_overlap must be non-negative and smaller than chunk_size") + self.path = str(path) + self.chunk_size = chunk_size + self.chunk_overlap = chunk_overlap + self._fts5 = True + Path(self.path).parent.mkdir(parents=True, exist_ok=True) + super().__init__([ + CapabilitySpec("rag.ingest", write=True, persistent=True, risk=CapabilityRisk.ISOLATED_WRITE), + CapabilitySpec("rag.remove", write=True, persistent=True, risk=CapabilityRisk.DESTRUCTIVE), + CapabilitySpec("rag.list_documents", read=True), + CapabilitySpec("rag.search", read=True), + CapabilitySpec("rag.read_chunk", read=True), + CapabilitySpec("rag.reindex", write=True, persistent=True, risk=CapabilityRisk.SENSITIVE), + ]) + self._initialize() + + def _connect(self) -> sqlite3.Connection: + connection = sqlite3.connect(self.path, timeout=30) + connection.row_factory = sqlite3.Row + return connection + + def _initialize(self) -> None: + with self._connect() as connection: + connection.execute( + """ + CREATE TABLE IF NOT EXISTS rag_documents ( + document_id TEXT PRIMARY KEY, + payload TEXT NOT NULL + ) + """ + ) + connection.execute( + """ + CREATE TABLE IF NOT EXISTS rag_chunks ( + chunk_id TEXT PRIMARY KEY, + document_id TEXT NOT NULL, + position INTEGER NOT NULL, + content TEXT NOT NULL, + FOREIGN KEY(document_id) REFERENCES rag_documents(document_id) ON DELETE CASCADE + ) + """ + ) + try: + connection.execute( + "CREATE VIRTUAL TABLE IF NOT EXISTS rag_chunks_fts USING fts5(chunk_id UNINDEXED, content)" + ) + except sqlite3.OperationalError: + self._fts5 = False + + async def health(self) -> ProviderHealth: + return ProviderHealth( + healthy=True, + status="ready" if self._fts5 else "degraded", + message=None if self._fts5 else "SQLite FTS5 unavailable; LIKE search fallback is active", + degraded_capabilities=[] if self._fts5 else ["rag.search"], + ) + + def ingest(self, document: RetrievalDocument | dict[str, Any]) -> str: + value = document if isinstance(document, RetrievalDocument) else RetrievalDocument(**document) + chunks = self._chunks(value.content) + with self._connect() as connection: + connection.execute("PRAGMA foreign_keys=ON") + connection.execute( + "INSERT OR REPLACE INTO rag_documents(document_id, payload) VALUES (?, ?)", + (value.document_id, json.dumps(value.to_dict(), ensure_ascii=False)), + ) + old_chunks = connection.execute( + "SELECT chunk_id FROM rag_chunks WHERE document_id = ?", (value.document_id,) + ).fetchall() + if self._fts5: + connection.executemany( + "DELETE FROM rag_chunks_fts WHERE chunk_id = ?", + [(row["chunk_id"],) for row in old_chunks], + ) + connection.execute("DELETE FROM rag_chunks WHERE document_id = ?", (value.document_id,)) + for position, content in enumerate(chunks): + chunk_id = f"{value.document_id}:{position}" + connection.execute( + "INSERT INTO rag_chunks(chunk_id, document_id, position, content) VALUES (?, ?, ?, ?)", + (chunk_id, value.document_id, position, content), + ) + if self._fts5: + connection.execute( + "INSERT INTO rag_chunks_fts(chunk_id, content) VALUES (?, ?)", + (chunk_id, content), + ) + return value.document_id + + def remove(self, document_id: str) -> bool: + with self._connect() as connection: + connection.execute("PRAGMA foreign_keys=ON") + rows = connection.execute( + "SELECT chunk_id FROM rag_chunks WHERE document_id = ?", (document_id,) + ).fetchall() + if self._fts5: + connection.executemany( + "DELETE FROM rag_chunks_fts WHERE chunk_id = ?", [(row["chunk_id"],) for row in rows] + ) + cursor = connection.execute("DELETE FROM rag_documents WHERE document_id = ?", (document_id,)) + return cursor.rowcount > 0 + + def list_documents(self) -> list[RetrievalDocument]: + with self._connect() as connection: + rows = connection.execute("SELECT payload FROM rag_documents ORDER BY rowid").fetchall() + return [RetrievalDocument(**json.loads(row["payload"])) for row in rows] + + def search( + self, + query: str, + *, + limit: int = 5, + scope: str | None = None, + owner_id: str | None = None, + tenant_id: str | None = None, + ) -> list[RetrievalResult]: + if not query.strip(): + return [] + if limit < 1: + raise ValueError("limit must be at least 1") + with self._connect() as connection: + if self._fts5: + rows = connection.execute( + """ + SELECT c.*, d.payload, bm25(rag_chunks_fts) AS rank + FROM rag_chunks_fts + JOIN rag_chunks c ON c.chunk_id = rag_chunks_fts.chunk_id + JOIN rag_documents d ON d.document_id = c.document_id + WHERE rag_chunks_fts MATCH ? + ORDER BY rank + LIMIT ? + """, + (query, max(limit * 10, limit)), + ).fetchall() + else: + rows = connection.execute( + """ + SELECT c.*, d.payload, NULL AS rank + FROM rag_chunks c + JOIN rag_documents d ON d.document_id = c.document_id + WHERE lower(c.content) LIKE ? + ORDER BY c.rowid + LIMIT ? + """, + (f"%{query.lower()}%", max(limit * 10, limit)), + ).fetchall() + results = [] + for row in rows: + document = RetrievalDocument(**json.loads(row["payload"])) + if scope is not None and document.scope != scope: + continue + if owner_id is not None and document.owner_id != owner_id: + continue + if tenant_id is not None and document.tenant_id != tenant_id: + continue + results.append(self._result(row, document)) + if len(results) >= limit: + break + return results + + def read_chunk(self, chunk_id: str) -> RetrievalResult | None: + with self._connect() as connection: + row = connection.execute( + """ + SELECT c.*, d.payload, NULL AS rank + FROM rag_chunks c JOIN rag_documents d ON d.document_id = c.document_id + WHERE c.chunk_id = ? + """, + (chunk_id,), + ).fetchone() + if row is None: + return None + return self._result(row, RetrievalDocument(**json.loads(row["payload"]))) + + def reindex(self) -> None: + if not self._fts5: + return + with self._connect() as connection: + connection.execute("DELETE FROM rag_chunks_fts") + rows = connection.execute("SELECT chunk_id, content FROM rag_chunks ORDER BY rowid").fetchall() + connection.executemany( + "INSERT INTO rag_chunks_fts(chunk_id, content) VALUES (?, ?)", + [(row["chunk_id"], row["content"]) for row in rows], + ) + + async def invoke(self, capability: str, **arguments: Any) -> Any: + method_name = capability.removeprefix("rag.") + method = getattr(self, method_name) + return method(**arguments) + + def _chunks(self, content: str) -> list[str]: + if not content: + return [""] + step = self.chunk_size - self.chunk_overlap + return [content[index:index + self.chunk_size] for index in range(0, len(content), step)] + + @staticmethod + def _result(row: sqlite3.Row, document: RetrievalDocument) -> RetrievalResult: + rank = row["rank"] + return RetrievalResult( + document_id=document.document_id, + chunk_id=row["chunk_id"], + content=row["content"], + citation_id=f"rag:{row['chunk_id']}", + title=document.title, + source=document.source, + position=row["position"], + score=None if rank is None else float(-rank), + scope=document.scope, + owner_id=document.owner_id, + tenant_id=document.tenant_id, + metadata=deepcopy(document.metadata), + ) + + +class SessionSearchProvider(BaseCapabilityProvider): + """Literal cross-Session search kept separate from knowledge-base RAG.""" + + name = "session-search" + version = "1" + + def __init__(self, store: SessionStore): + self.store = store + super().__init__([CapabilitySpec("session.search", read=True)]) + + def search(self, query: str, *, limit: int = 20) -> list[dict[str, Any]]: + if not query.strip(): + return [] + needle = query.casefold() + matches: list[dict[str, Any]] = [] + for session in self.store.list(limit=max(limit, 100)): + for event in reversed(session.events): + rendered = json.dumps(event.data, ensure_ascii=False, sort_keys=True) + if needle not in rendered.casefold(): + continue + matches.append({ + "session_id": session.session_id, + "event_id": event.event_id, + "sequence": event.sequence, + "event_type": event.type, + "citation_id": f"session:{session.session_id}:{event.sequence}", + "timestamp": event.timestamp, + "data": deepcopy(event.data), + }) + if len(matches) >= limit: + return matches + return matches + + async def invoke(self, capability: str, **arguments: Any) -> Any: + if capability != "session.search": + raise LookupError(capability) + return self.search(**arguments) + + +class SkillProviderAdapter(BaseCapabilityProvider): + name = "skills" + version = "1" + + def __init__(self, skill_manager: Any): + self.skill_manager = skill_manager + super().__init__([ + CapabilitySpec("skill.list", read=True), + CapabilitySpec("skill.activate", read=True), + CapabilitySpec("skill.read_reference", read=True), + CapabilitySpec( + "skill.execute_script", + execute=True, + risk=CapabilityRisk.SENSITIVE, + requires_sandbox=True, + requires_approval=True, + ), + ]) + + async def invoke(self, capability: str, **arguments: Any) -> Any: + if capability == "skill.list": + return [asdict(skill) for skill in self.skill_manager.skills.values()] + method_name = capability.removeprefix("skill.") + if method_name == "activate": + method_name = "activate_skill" + method = getattr(self.skill_manager, method_name) + return method(**arguments) + + +class MCPProviderAdapter(BaseCapabilityProvider): + name = "mcp" + version = "1" + + def __init__(self, manager: Any, tool_names: Iterable[str] | None = None): + self.manager = manager + names = list(tool_names or []) + super().__init__([ + CapabilitySpec( + f"mcp.{name}", + network=True, + execute=True, + risk=CapabilityRisk.SENSITIVE, + cancellable=True, + ) + for name in names + ]) + + async def invoke(self, capability: str, **arguments: Any) -> Any: + tool_name = capability.removeprefix("mcp.") + return await self.manager.call_tool(tool_name, arguments) + + +class WorkflowProviderAdapter(BaseCapabilityProvider): + name = "lightflow" + version = "1" + + def __init__(self, flow: Any): + self.flow = flow + super().__init__([ + CapabilitySpec("workflow.start", execute=True, cancellable=True), + CapabilitySpec("workflow.status", read=True), + CapabilitySpec("workflow.resume", execute=True, resumable=True), + CapabilitySpec("workflow.rerun_step", execute=True, risk=CapabilityRisk.SENSITIVE), + ]) + + async def invoke(self, capability: str, **arguments: Any) -> Any: + if capability == "workflow.start": + method = getattr(self.flow, "arun", None) + if callable(method): + return await method(**arguments) + return await _call_in_thread(self.flow.run, **arguments) + if capability == "workflow.status": + return self.flow.get_run(arguments["run_id"]) + if capability == "workflow.resume": + return await _call_in_thread(self.flow.resume, arguments["run_id"], **{ + key: value for key, value in arguments.items() if key != "run_id" + }) + if capability == "workflow.rerun_step": + return await _call_in_thread( + self.flow.rerun_step, + arguments["run_id"], + arguments["step_name"], + ) + raise LookupError(capability) + + +async def _call_in_thread(function: Any, *args: Any, **kwargs: Any) -> Any: + result = await __import__("asyncio").to_thread(function, *args, **kwargs) + if inspect.isawaitable(result): + return await result + return result + + +__all__ = [ + "RetrievalDocument", + "RetrievalResult", + "RetrievalProvider", + "SqliteFTSRetrievalProvider", + "SessionSearchProvider", + "SkillProviderAdapter", + "MCPProviderAdapter", + "WorkflowProviderAdapter", +] diff --git a/LightAgent/mcp_client_manager.py b/LightAgent/mcp_client_manager.py index e53818e..9512537 100644 --- a/LightAgent/mcp_client_manager.py +++ b/LightAgent/mcp_client_manager.py @@ -7,12 +7,17 @@ """ import re +import inspect from functools import partial from typing import Optional, Dict, Any from mcp import ClientSession, StdioServerParameters from mcp.client.sse import sse_client from mcp.client.stdio import stdio_client +try: + from mcp.client.streamable_http import streamablehttp_client +except ImportError: # Older MCP SDKs retain stdio/SSE compatibility. + streamablehttp_client = None from contextlib import AsyncExitStack from .tools import ToolRegistry # 关键修改:从当前包导入 @@ -21,21 +26,39 @@ class MCPClientManager: """增强版MCP客户端管理器""" - def __init__(self, config: dict, tool_registry: ToolRegistry): + def __init__(self, config: dict, tool_registry: ToolRegistry, credential_provider: Any = None): self.config = config self.tool_registry = tool_registry self.session: Optional[ClientSession] = None self.exit_stack = AsyncExitStack() self.server_sessions = {} self.last_mcp_errors = [] + self.credential_provider = credential_provider + self._registered_tools: Dict[str, set[str]] = {} + self._tool_owners: Dict[str, str] = {} async def _create_session(self, server_name: str, config: dict): """创建并管理会话上下文""" - if 'url' in config: + headers = dict(config.get('headers', {})) + if self.credential_provider is not None: + resolver = getattr(self.credential_provider, "headers", self.credential_provider) + supplied = resolver(server_name, dict(config)) + if inspect.isawaitable(supplied): + supplied = await supplied + headers.update(dict(supplied or {})) + transport_name = str(config.get("transport", "")).lower() + if transport_name in {"streamable-http", "streamable_http", "http"}: + if streamablehttp_client is None: + raise RuntimeError("installed MCP SDK does not support Streamable HTTP") + streams_context = streamablehttp_client(url=config['url'], headers=headers) + streams = await self.exit_stack.enter_async_context(streams_context) + session_context = ClientSession(*streams[:2]) + self.session = await self.exit_stack.enter_async_context(session_context) + elif 'url' in config: # SSE 服务器连接 streams_context = sse_client( url=config['url'], - headers=config.get('headers', {}) + headers=headers ) streams = await self.exit_stack.enter_async_context(streams_context) session_context = ClientSession(*streams) @@ -59,6 +82,8 @@ async def cleanup(self): """清理所有会话资源""" await self.exit_stack.aclose() self.server_sessions.clear() + self.session = None + self.exit_stack = AsyncExitStack() async def register_mcp_tool(self) -> bool: """自动注册所有MCP服务的工具""" @@ -67,7 +92,7 @@ async def register_mcp_tool(self) -> bool: enabled_servers = [ (name, config) for name, config in self.config["mcpServers"].items() - if not config["disabled"] + if not config.get("disabled", False) ] for server_name, config in enabled_servers: @@ -78,9 +103,10 @@ async def register_mcp_tool(self) -> bool: for tool in tools_response.tools: try: + public_name = self._public_tool_name(server_name, config, tool.name) # 构建工具元数据 tool_info = { - "tool_name": tool.name, + "tool_name": public_name, "tool_description": tool.description, "tool_params": [] } @@ -98,8 +124,8 @@ async def register_mcp_tool(self) -> bool: }) # 注册到工具注册表 - self.tool_registry.function_info[tool.name] = tool_info - self.tool_registry.function_mappings[tool.name] = partial( + self.tool_registry.function_info[public_name] = tool_info + self.tool_registry.function_mappings[public_name] = partial( self._call_tool_wrapper, tool_name=tool.name, target_server=server_name @@ -109,7 +135,7 @@ async def register_mcp_tool(self) -> bool: openai_schema = { "type": "function", "function": { - "name": tool.name, + "name": public_name, "description": tool.description, "parameters": { "type": "object", @@ -122,6 +148,8 @@ async def register_mcp_tool(self) -> bool: } } self.tool_registry.openai_function_schemas.append(openai_schema) + self._registered_tools.setdefault(server_name, set()).add(public_name) + self._tool_owners[public_name] = server_name registered_count += 1 print(f"✅ The registered MCP tool : {tool.name}") except Exception as e: @@ -134,6 +162,40 @@ async def register_mcp_tool(self) -> bool: await self.cleanup() return registered_count > 0 + async def refresh_tools(self, server_name: str | None = None) -> bool: + """Refresh MCP tool lists without leaving duplicate registrations.""" + targets = [server_name] if server_name else list(self._registered_tools) + for target in targets: + for name in self._registered_tools.pop(target, set()): + self.tool_registry.function_info.pop(name, None) + self.tool_registry.function_mappings.pop(name, None) + self._tool_owners.pop(name, None) + self.tool_registry.openai_function_schemas = [ + schema for schema in self.tool_registry.openai_function_schemas + if schema.get("function", {}).get("name") != name + ] + if server_name is None: + return await self.register_mcp_tool() + original = self.config["mcpServers"] + self.config["mcpServers"] = { + name: {**config, "disabled": name != server_name} + for name, config in original.items() + } + try: + return await self.register_mcp_tool() + finally: + self.config["mcpServers"] = original + + def _public_tool_name(self, server_name: str, config: dict, tool_name: str) -> str: + namespace = re.sub(r"[^A-Za-z0-9_-]", "_", str(config.get("namespace") or server_name)) + namespaced = f"{namespace}__{tool_name}" + if config.get("namespace_tools") or ( + tool_name in self.tool_registry.function_mappings + and self._tool_owners.get(tool_name) != server_name + ): + return namespaced + return tool_name + async def _call_tool_wrapper(self, tool_name: str, target_server: str, **kwargs): """参数转换适配器""" return await self.call_tool( @@ -148,38 +210,42 @@ async def call_tool(self, tool_name: str, arguments: dict, target_server: str = enabled_servers = [ (name, config) for name, config in self.config["mcpServers"].items() - if not config["disabled"] + if not config.get("disabled", False) ] if target_server: enabled_servers = [s for s in enabled_servers if s[0] == target_server] for server_name, config in enabled_servers: - try: - session = self.server_sessions.get(server_name) - if not session: - await self._create_session(server_name, config) - session = self.session + max_attempts = max(1, int(config.get("reconnect_attempts", 0)) + 1) + for attempt in range(max_attempts): + try: + session = self.server_sessions.get(server_name) + if not session: + await self._create_session(server_name, config) + session = self.session - tools = await session.list_tools() - available_tools = {t.name: t for t in tools.tools} + tools = await session.list_tools() + available_tools = {t.name: t for t in tools.tools} - if tool_name in available_tools: - # 验证参数类型 - schema = available_tools[tool_name].inputSchema - self._validate_arguments(arguments, schema) + if tool_name in available_tools: + # 验证参数类型 + schema = available_tools[tool_name].inputSchema + self._validate_arguments(arguments, schema) - # 执行调用 - result = await session.call_tool(tool_name, arguments) + # 执行调用 + result = await session.call_tool(tool_name, arguments) + await self.cleanup() + return { + "server": server_name, + "tool": tool_name, + "result": result.content[0].text + } + except Exception as e: + self._record_error(server_name, f"call_tool:{tool_name}", e) await self.cleanup() - return { - "server": server_name, - "tool": tool_name, - "result": result.content[0].text - } - except Exception as e: - self._record_error(server_name, f"call_tool:{tool_name}", e) - continue + if attempt + 1 >= max_attempts: + break raise ValueError(self._format_tool_not_found(tool_name)) diff --git a/LightAgent/runtime.py b/LightAgent/runtime.py new file mode 100644 index 0000000..f4db7ca --- /dev/null +++ b/LightAgent/runtime.py @@ -0,0 +1,843 @@ +"""Long-task runtime controls built on the Session event log.""" + +from __future__ import annotations + +import asyncio +import inspect +from copy import deepcopy +from dataclasses import asdict, dataclass, field +from enum import Enum +from typing import Any, Awaitable, Callable, Iterable +from uuid import uuid4 + +from .capabilities import CapabilityRegistry, PermissionSet, RuntimeContext +from .session import InMemorySessionStore, Session, SessionCheckpoint, SessionStore, _utc_now + + +class InboxMessageType(str, Enum): + FOLLOWUP = "followup" + STEERING = "steering" + CONTEXT = "context" + APPROVAL = "approval" + + +class InboxMessageStatus(str, Enum): + PENDING = "pending" + CLAIMED = "claimed" + COMPLETED = "completed" + REJECTED = "rejected" + + +@dataclass +class InboxMessage: + type: InboxMessageType + content: Any + message_id: str = field(default_factory=lambda: uuid4().hex) + status: InboxMessageStatus = InboxMessageStatus.PENDING + created_at: str = field(default_factory=_utc_now) + claimed_at: str | None = None + completed_at: str | None = None + correlation_id: str | None = None + metadata: dict[str, Any] = field(default_factory=dict) + + def to_dict(self) -> dict[str, Any]: + value = asdict(self) + value["type"] = self.type.value + value["status"] = self.status.value + return value + + @classmethod + def from_dict(cls, value: dict[str, Any]) -> "InboxMessage": + return cls( + message_id=str(value.get("message_id") or uuid4().hex), + type=InboxMessageType(value["type"]), + content=deepcopy(value.get("content")), + status=InboxMessageStatus(value.get("status", "pending")), + created_at=str(value.get("created_at") or _utc_now()), + claimed_at=value.get("claimed_at"), + completed_at=value.get("completed_at"), + correlation_id=value.get("correlation_id"), + metadata=dict(value.get("metadata") or {}), + ) + + +class AgentInbox: + """Ordered, idempotent Inbox whose mutations are Session events.""" + + def __init__(self, event_sink: Callable[[str, dict[str, Any]], Any] | None = None): + self._messages: list[InboxMessage] = [] + self._event_sink = event_sink + + def enqueue( + self, + message_type: InboxMessageType | str, + content: Any, + *, + message_id: str | None = None, + correlation_id: str | None = None, + metadata: dict[str, Any] | None = None, + ) -> InboxMessage: + resolved_id = message_id or uuid4().hex + existing = self.get(resolved_id) + if existing is not None: + return existing + message = InboxMessage( + message_id=resolved_id, + type=InboxMessageType(message_type), + content=deepcopy(content), + correlation_id=correlation_id, + metadata=deepcopy(metadata or {}), + ) + self._messages.append(message) + self._emit("inbox.enqueued", {"message": message.to_dict()}) + return deepcopy(message) + + def pending(self, message_type: InboxMessageType | str | None = None) -> list[InboxMessage]: + resolved_type = InboxMessageType(message_type) if message_type is not None else None + return deepcopy([ + message for message in self._messages + if message.status == InboxMessageStatus.PENDING + and (resolved_type is None or message.type == resolved_type) + ]) + + def claim_next(self, *, safe_boundary: bool = True) -> InboxMessage | None: + for message in self._messages: + if message.status != InboxMessageStatus.PENDING: + continue + if message.type == InboxMessageType.STEERING and not safe_boundary: + continue + message.status = InboxMessageStatus.CLAIMED + message.claimed_at = _utc_now() + self._emit("inbox.claimed", {"message": message.to_dict()}) + return deepcopy(message) + return None + + def complete(self, message_id: str, *, result: Any = None) -> InboxMessage: + message = self._require(message_id) + if message.status == InboxMessageStatus.COMPLETED: + return deepcopy(message) + if message.status != InboxMessageStatus.CLAIMED: + raise ValueError("only claimed Inbox messages can be completed") + message.status = InboxMessageStatus.COMPLETED + message.completed_at = _utc_now() + self._emit("inbox.completed", {"message": message.to_dict(), "result": deepcopy(result)}) + return deepcopy(message) + + def reject(self, message_id: str, reason: str) -> InboxMessage: + message = self._require(message_id) + message.status = InboxMessageStatus.REJECTED + message.completed_at = _utc_now() + self._emit("inbox.rejected", {"message": message.to_dict(), "reason": reason}) + return deepcopy(message) + + def get(self, message_id: str) -> InboxMessage | None: + message = next((item for item in self._messages if item.message_id == message_id), None) + return deepcopy(message) if message else None + + def list(self) -> list[InboxMessage]: + return deepcopy(self._messages) + + def restore(self, session: Session) -> None: + self._messages = [] + by_id: dict[str, InboxMessage] = {} + for event in session.events: + if not event.type.startswith("inbox."): + continue + value = event.data.get("message") + if not isinstance(value, dict): + continue + message = InboxMessage.from_dict(value) + by_id[message.message_id] = message + self._messages = sorted(by_id.values(), key=lambda item: item.created_at) + + def _require(self, message_id: str) -> InboxMessage: + message = next((item for item in self._messages if item.message_id == message_id), None) + if message is None: + raise KeyError(message_id) + return message + + def _emit(self, event_type: str, data: dict[str, Any]) -> None: + if self._event_sink: + self._event_sink(event_type, data) + + +class GoalStatus(str, Enum): + PENDING = "pending" + ACTIVE = "active" + COMPLETED = "completed" + BLOCKED = "blocked" + CANCELLED = "cancelled" + + +@dataclass +class Goal: + objective: str + goal_id: str = field(default_factory=lambda: uuid4().hex) + status: GoalStatus = GoalStatus.PENDING + acceptance_criteria: list[str] = field(default_factory=list) + parent_goal_id: str | None = None + evidence: list[Any] = field(default_factory=list) + blocker: str | None = None + created_at: str = field(default_factory=_utc_now) + updated_at: str = field(default_factory=_utc_now) + metadata: dict[str, Any] = field(default_factory=dict) + + def to_dict(self) -> dict[str, Any]: + value = asdict(self) + value["status"] = self.status.value + return value + + @classmethod + def from_dict(cls, value: dict[str, Any]) -> "Goal": + return cls( + goal_id=str(value.get("goal_id") or uuid4().hex), + objective=str(value["objective"]), + status=GoalStatus(value.get("status", "pending")), + acceptance_criteria=list(value.get("acceptance_criteria") or []), + parent_goal_id=value.get("parent_goal_id"), + evidence=deepcopy(value.get("evidence") or []), + blocker=value.get("blocker"), + created_at=str(value.get("created_at") or _utc_now()), + updated_at=str(value.get("updated_at") or _utc_now()), + metadata=dict(value.get("metadata") or {}), + ) + + +class GoalManager: + def __init__(self, event_sink: Callable[[str, dict[str, Any]], Any] | None = None): + self._goals: dict[str, Goal] = {} + self._event_sink = event_sink + + def create( + self, + objective: str, + *, + acceptance_criteria: Iterable[str] | None = None, + parent_goal_id: str | None = None, + metadata: dict[str, Any] | None = None, + ) -> Goal: + if not objective.strip(): + raise ValueError("goal objective must not be empty") + if parent_goal_id and parent_goal_id not in self._goals: + raise KeyError(parent_goal_id) + goal = Goal( + objective=objective, + acceptance_criteria=list(acceptance_criteria or []), + parent_goal_id=parent_goal_id, + metadata=deepcopy(metadata or {}), + ) + self._goals[goal.goal_id] = goal + self._emit("goal.created", goal) + return deepcopy(goal) + + def activate(self, goal_id: str) -> Goal: + return self._transition(goal_id, GoalStatus.ACTIVE) + + def complete(self, goal_id: str, *, evidence: Iterable[Any] | None = None) -> Goal: + goal = self._require(goal_id) + goal.evidence.extend(deepcopy(list(evidence or []))) + goal.blocker = None + return self._transition(goal_id, GoalStatus.COMPLETED, event_type="goal.completed") + + def block(self, goal_id: str, reason: str) -> Goal: + goal = self._require(goal_id) + goal.blocker = reason + return self._transition(goal_id, GoalStatus.BLOCKED, event_type="goal.blocked") + + def cancel(self, goal_id: str, reason: str | None = None) -> Goal: + goal = self._require(goal_id) + goal.blocker = reason + return self._transition(goal_id, GoalStatus.CANCELLED, event_type="goal.cancelled") + + def get(self, goal_id: str) -> Goal: + return deepcopy(self._require(goal_id)) + + def list(self, *, status: GoalStatus | str | None = None) -> list[Goal]: + resolved = GoalStatus(status) if status is not None else None + return deepcopy([goal for goal in self._goals.values() if resolved is None or goal.status == resolved]) + + def restore(self, session: Session) -> None: + by_id: dict[str, Goal] = {} + for event in session.events: + if not event.type.startswith("goal."): + continue + value = event.data.get("goal") + if isinstance(value, dict): + goal = Goal.from_dict(value) + by_id[goal.goal_id] = goal + self._goals = by_id + + def _transition(self, goal_id: str, status: GoalStatus, event_type: str = "goal.updated") -> Goal: + goal = self._require(goal_id) + if goal.status in {GoalStatus.COMPLETED, GoalStatus.CANCELLED} and goal.status != status: + raise ValueError(f"terminal goal `{goal_id}` cannot transition from {goal.status.value}") + goal.status = status + goal.updated_at = _utc_now() + self._emit(event_type, goal) + return deepcopy(goal) + + def _require(self, goal_id: str) -> Goal: + if goal_id not in self._goals: + raise KeyError(goal_id) + return self._goals[goal_id] + + def _emit(self, event_type: str, goal: Goal) -> None: + if self._event_sink: + self._event_sink(event_type, {"goal": goal.to_dict()}) + + +@dataclass(frozen=True) +class BudgetLimits: + model_calls: int | None = None + tool_calls: int | None = None + tokens: int | None = None + seconds: float | None = None + cost: float | None = None + + def __post_init__(self) -> None: + for name, value in asdict(self).items(): + if value is not None and value < 0: + raise ValueError(f"budget limit `{name}` must be non-negative") + + +@dataclass +class BudgetUsage: + model_calls: int = 0 + tool_calls: int = 0 + tokens: int = 0 + seconds: float = 0.0 + cost: float = 0.0 + + def to_dict(self) -> dict[str, Any]: + return asdict(self) + + +class BudgetExceeded(RuntimeError): + def __init__(self, dimension: str, limit: float, used: float): + super().__init__(f"budget exceeded for {dimension}: used={used}, limit={limit}") + self.dimension = dimension + self.limit = limit + self.used = used + + +class BudgetManager: + def __init__( + self, + limits: BudgetLimits | None = None, + event_sink: Callable[[str, dict[str, Any]], Any] | None = None, + ): + self.limits = limits or BudgetLimits() + self.usage = BudgetUsage() + self._event_sink = event_sink + + def consume( + self, + *, + model_calls: int = 0, + tool_calls: int = 0, + tokens: int = 0, + seconds: float = 0.0, + cost: float = 0.0, + ) -> BudgetUsage: + candidate = BudgetUsage( + model_calls=self.usage.model_calls + model_calls, + tool_calls=self.usage.tool_calls + tool_calls, + tokens=self.usage.tokens + tokens, + seconds=self.usage.seconds + seconds, + cost=self.usage.cost + cost, + ) + for dimension, limit in asdict(self.limits).items(): + used = getattr(candidate, dimension) + if limit is not None and used > limit: + if self._event_sink: + self._event_sink("budget.exhausted", { + "dimension": dimension, + "limit": limit, + "used": used, + }) + raise BudgetExceeded(dimension, limit, used) + self.usage = candidate + if self._event_sink: + self._event_sink("budget.consumed", {"usage": self.usage.to_dict()}) + return deepcopy(self.usage) + + def remaining(self) -> dict[str, float | int | None]: + return { + dimension: None if limit is None else max(0, limit - getattr(self.usage, dimension)) + for dimension, limit in asdict(self.limits).items() + } + + def restore(self, session: Session) -> None: + for event in reversed(session.events): + if event.type == "budget.consumed" and isinstance(event.data.get("usage"), dict): + self.usage = BudgetUsage(**event.data["usage"]) + return + self.usage = BudgetUsage() + + +@dataclass +class ProgressState: + steps: int = 0 + unchanged_steps: int = 0 + repeated_tool_calls: int = 0 + last_output: str | None = None + last_tool_signature: str | None = None + + +class ProgressTracker: + """Detect bounded no-progress and repeated-tool loops.""" + + def __init__( + self, + *, + max_unchanged_steps: int = 3, + max_repeated_tool_calls: int = 3, + event_sink: Callable[[str, dict[str, Any]], Any] | None = None, + ): + self.max_unchanged_steps = max_unchanged_steps + self.max_repeated_tool_calls = max_repeated_tool_calls + self.state = ProgressState() + self._event_sink = event_sink + + def record(self, *, output: Any = None, tool: str | None = None, arguments: Any = None) -> ProgressState: + rendered = repr(output) + signature = repr((tool, arguments)) if tool else None + self.state.steps += 1 + self.state.unchanged_steps = self.state.unchanged_steps + 1 if rendered == self.state.last_output else 0 + self.state.repeated_tool_calls = ( + self.state.repeated_tool_calls + 1 + if signature is not None and signature == self.state.last_tool_signature + else 0 + ) + self.state.last_output = rendered + self.state.last_tool_signature = signature + if self.state.unchanged_steps >= self.max_unchanged_steps: + self._emit("progress.stalled", "unchanged_output") + if self.state.repeated_tool_calls >= self.max_repeated_tool_calls: + self._emit("progress.stalled", "repeated_tool_call") + return deepcopy(self.state) + + @property + def stalled(self) -> bool: + return ( + self.state.unchanged_steps >= self.max_unchanged_steps + or self.state.repeated_tool_calls >= self.max_repeated_tool_calls + ) + + def _emit(self, event_type: str, reason: str) -> None: + if self._event_sink: + self._event_sink(event_type, {"reason": reason, "state": asdict(self.state)}) + + +class JobStatus(str, Enum): + PENDING = "pending" + RUNNING = "running" + SUCCESS = "success" + FAILED = "failed" + CANCELLED = "cancelled" + INTERRUPTED = "interrupted" + + +@dataclass +class JobRecord: + name: str + job_id: str = field(default_factory=lambda: uuid4().hex) + status: JobStatus = JobStatus.PENDING + owner_agent_id: str | None = None + created_at: str = field(default_factory=_utc_now) + started_at: str | None = None + completed_at: str | None = None + result: Any = None + error: str | None = None + metadata: dict[str, Any] = field(default_factory=dict) + output: list[Any] = field(default_factory=list) + + def to_dict(self) -> dict[str, Any]: + value = asdict(self) + value["status"] = self.status.value + return value + + +class JobManager: + """Manage cancellable background coroutines and persist their state.""" + + def __init__( + self, + event_sink: Callable[[str, dict[str, Any]], Any] | None = None, + inbox: AgentInbox | None = None, + ): + self._records: dict[str, JobRecord] = {} + self._tasks: dict[str, asyncio.Task[Any]] = {} + self._event_sink = event_sink + self._inbox = inbox + + def start( + self, + name: str, + operation: Callable[[], Any] | Awaitable[Any], + *, + owner_agent_id: str | None = None, + metadata: dict[str, Any] | None = None, + ) -> JobRecord: + loop = asyncio.get_running_loop() + record = JobRecord(name=name, owner_agent_id=owner_agent_id, metadata=deepcopy(metadata or {})) + self._records[record.job_id] = record + self._emit("job.created", record) + self._tasks[record.job_id] = loop.create_task(self._execute(record.job_id, operation)) + return deepcopy(record) + + async def _execute(self, job_id: str, operation: Callable[[], Any] | Awaitable[Any]) -> None: + record = self._records[job_id] + record.status = JobStatus.RUNNING + record.started_at = _utc_now() + self._emit("job.started", record) + try: + value = operation() if callable(operation) else operation + if inspect.isawaitable(value): + value = await value + record.result = value + record.status = JobStatus.SUCCESS + record.completed_at = _utc_now() + self._emit("job.completed", record) + except asyncio.CancelledError: + record.status = JobStatus.CANCELLED + record.error = "cancelled" + record.completed_at = _utc_now() + self._emit("job.cancelled", record) + raise + except Exception as error: + record.status = JobStatus.FAILED + record.error = f"{type(error).__name__}: {error}" + record.completed_at = _utc_now() + self._emit("job.failed", record) + finally: + if record.completed_at is None: + record.completed_at = _utc_now() + if self._inbox: + self._inbox.enqueue( + InboxMessageType.CONTEXT, + {"job": record.to_dict()}, + correlation_id=record.job_id, + metadata={"kind": "job_completion"}, + ) + + async def wait(self, job_id: str) -> JobRecord: + if job_id not in self._records: + raise KeyError(job_id) + task = self._tasks.get(job_id) + if task is not None: + try: + await task + except asyncio.CancelledError: + pass + return deepcopy(self._records[job_id]) + + def cancel(self, job_id: str) -> bool: + task = self._tasks.get(job_id) + if task is None or task.done(): + return False + return task.cancel() + + def emit_output(self, job_id: str, value: Any) -> JobRecord: + record = self._records.get(job_id) + if record is None: + raise KeyError(job_id) + record.output.append(deepcopy(value)) + self._emit("job.output", record) + return deepcopy(record) + + def get(self, job_id: str) -> JobRecord: + if job_id not in self._records: + raise KeyError(job_id) + return deepcopy(self._records[job_id]) + + def list(self) -> list[JobRecord]: + return deepcopy(list(self._records.values())) + + def mark_interrupted(self) -> None: + for record in self._records.values(): + if record.status in {JobStatus.PENDING, JobStatus.RUNNING}: + record.status = JobStatus.INTERRUPTED + record.completed_at = _utc_now() + self._emit("job.interrupted", record) + + def restore(self, session: Session) -> None: + records: dict[str, JobRecord] = {} + for event in session.events: + if not event.type.startswith("job."): + continue + value = event.data.get("job") + if not isinstance(value, dict): + continue + payload = dict(value) + payload["status"] = JobStatus(payload.get("status", "pending")) + record = JobRecord(**payload) + records[record.job_id] = record + self._records = records + self._tasks = {} + + def _emit(self, event_type: str, record: JobRecord) -> None: + if self._event_sink: + self._event_sink(event_type, {"job": record.to_dict()}) + + +@dataclass +class SubagentRecord: + agent_id: str + name: str + parent_agent_id: str | None + depth: int + permissions: PermissionSet + persistent: bool = False + status: str = "ready" + metadata: dict[str, Any] = field(default_factory=dict) + + +class SubagentManager: + def __init__( + self, + *, + max_depth: int = 2, + max_agents: int = 8, + max_concurrency: int = 4, + event_sink: Callable[[str, dict[str, Any]], Any] | None = None, + ): + self.max_depth = max_depth + self.max_agents = max_agents + self.max_concurrency = max_concurrency + self._event_sink = event_sink + self._agents: dict[str, tuple[Any, SubagentRecord]] = {} + self._running = 0 + + def register( + self, + agent: Any, + *, + parent_agent_id: str | None = None, + parent_permissions: PermissionSet | None = None, + allowed_capabilities: Iterable[str] | None = None, + max_risk: Any = None, + persistent: bool = False, + metadata: dict[str, Any] | None = None, + ) -> SubagentRecord: + if len(self._agents) >= self.max_agents: + raise RuntimeError("subagent limit reached") + parent_record = self._agents.get(parent_agent_id, (None, None))[1] if parent_agent_id else None + depth = parent_record.depth + 1 if parent_record else 1 + if depth > self.max_depth: + raise RuntimeError("subagent depth limit reached") + base_permissions = parent_permissions or (parent_record.permissions if parent_record else PermissionSet()) + permissions = base_permissions.narrow( + allowed=allowed_capabilities, + max_risk=max_risk, + ) + agent_id = uuid4().hex + record = SubagentRecord( + agent_id=agent_id, + name=str(getattr(agent, "name", agent_id)), + parent_agent_id=parent_agent_id, + depth=depth, + permissions=permissions, + persistent=persistent, + metadata=deepcopy(metadata or {}), + ) + self._agents[agent_id] = (agent, record) + self._emit("subagent.created", record) + return deepcopy(record) + + async def run(self, agent_id: str, query: str, **kwargs: Any) -> Any: + if agent_id not in self._agents: + raise KeyError(agent_id) + agent, record = self._agents[agent_id] + if self._running >= self.max_concurrency: + raise RuntimeError("subagent concurrency limit reached") + self._running += 1 + record.status = "running" + self._emit("subagent.started", record) + try: + arun = getattr(agent, "arun", None) + if callable(arun): + result = await arun(query, **kwargs) + else: + result = await asyncio.to_thread(agent.run, query, **kwargs) + record.status = "success" + self._emit("subagent.completed", record, result=result) + return result + except Exception as error: + record.status = "failed" + self._emit("subagent.failed", record, error=f"{type(error).__name__}: {error}") + raise + finally: + self._running -= 1 + + def list(self) -> list[SubagentRecord]: + return deepcopy([record for _, record in self._agents.values()]) + + def tree(self) -> list[dict[str, Any]]: + return [ + { + "agent_id": record.agent_id, + "name": record.name, + "parent_agent_id": record.parent_agent_id, + "depth": record.depth, + "status": record.status, + "persistent": record.persistent, + } + for _, record in self._agents.values() + ] + + def send(self, agent_id: str, content: Any, *, message_type: InboxMessageType | str = InboxMessageType.CONTEXT) -> Any: + if agent_id not in self._agents: + raise KeyError(agent_id) + agent, _ = self._agents[agent_id] + runtime = getattr(agent, "runtime", None) + if runtime is None or runtime.session is None: + raise RuntimeError("target subagent does not have an open runtime Session") + return runtime.inbox.enqueue(message_type, content) + + def _emit(self, event_type: str, record: SubagentRecord, **extra: Any) -> None: + if self._event_sink: + data = { + "agent": { + **asdict(record), + "permissions": { + "allowed": sorted(record.permissions.allowed), + "denied": sorted(record.permissions.denied), + "max_risk": record.permissions.max_risk.value, + }, + }, + **extra, + } + self._event_sink(event_type, data) + + +class AgentRuntime: + """Bundle Session, Registry, Inbox, Goal, Budget, Job, and Subagent state.""" + + def __init__( + self, + *, + session_store: SessionStore | None = None, + capability_registry: CapabilityRegistry | None = None, + budget_limits: BudgetLimits | None = None, + runtime_id: str | None = None, + ): + self.runtime_id = runtime_id or uuid4().hex + self.session_store = session_store or InMemorySessionStore() + self.registry = capability_registry or CapabilityRegistry() + self.session: Session | None = None + self.context = RuntimeContext(runtime_id=self.runtime_id) + self.inbox = AgentInbox(self._append) + self.goals = GoalManager(self._append) + self.budget = BudgetManager(budget_limits, self._append) + self.jobs = JobManager(self._append, self.inbox) + self.subagents = SubagentManager(event_sink=self._append) + self.progress = ProgressTracker(event_sink=self._append) + + def open_session( + self, + session_id: str | None = None, + *, + metadata: dict[str, Any] | None = None, + agent_id: str | None = None, + user_id: str | None = None, + ) -> Session: + session = self.session_store.get(session_id) if session_id else None + if session is None: + session = Session(session_id=session_id or uuid4().hex, metadata=metadata or {}) + session.append("session.started", {"runtime_id": self.runtime_id}) + self.session_store.create(session) + self.session = session + self.context.session_id = session.session_id + self.context.agent_id = agent_id + self.context.user_id = user_id + self.inbox.restore(session) + self.goals.restore(session) + self.budget.restore(session) + self.jobs.restore(session) + self.jobs.mark_interrupted() + return deepcopy(session) + + def checkpoint(self, label: str | None = None, metadata: dict[str, Any] | None = None) -> SessionCheckpoint: + session = self._require_session() + checkpoint = session.checkpoint(label, metadata) + self.session_store.save(session) + return checkpoint + + def fork( + self, + *, + through_sequence: int | None = None, + metadata: dict[str, Any] | None = None, + ) -> Session: + forked = self._require_session().fork(through_sequence=through_sequence, metadata=metadata) + self.session_store.create(forked) + return forked + + def snapshot(self) -> dict[str, Any]: + session = self._require_session() + return { + "runtime_id": self.runtime_id, + "session": session.to_dict(), + "replay": session.replay().to_dict(), + "inbox": [message.to_dict() for message in self.inbox.list()], + "goals": [goal.to_dict() for goal in self.goals.list()], + "budget": { + "limits": asdict(self.budget.limits), + "usage": self.budget.usage.to_dict(), + "remaining": self.budget.remaining(), + }, + "jobs": [job.to_dict() for job in self.jobs.list()], + "progress": asdict(self.progress.state), + } + + def pause(self, reason: str | None = None) -> None: + self._append("session.paused", {"reason": reason}) + + def resume(self, reason: str | None = None) -> None: + self._append("session.resumed", {"reason": reason}) + + def cancel(self, reason: str | None = None) -> None: + self._append("session.cancelled", {"reason": reason}) + + def continue_run(self, reason: str | None = None) -> None: + self._append("session.continued", {"reason": reason}) + + def _append(self, event_type: str, data: dict[str, Any]) -> None: + session = self._require_session() + session.append( + event_type, + data, + turn_id=self.context.turn_id, + run_id=self.context.run_id, + agent_id=self.context.agent_id, + ) + self.session_store.save(session) + + def _require_session(self) -> Session: + if self.session is None: + raise RuntimeError("open_session() must be called first") + return self.session + + +__all__ = [ + "InboxMessageType", + "InboxMessageStatus", + "InboxMessage", + "AgentInbox", + "GoalStatus", + "Goal", + "GoalManager", + "BudgetLimits", + "BudgetUsage", + "BudgetExceeded", + "BudgetManager", + "ProgressState", + "ProgressTracker", + "JobStatus", + "JobRecord", + "JobManager", + "SubagentRecord", + "SubagentManager", + "AgentRuntime", +] diff --git a/LightAgent/session.py b/LightAgent/session.py new file mode 100644 index 0000000..58bdd2e --- /dev/null +++ b/LightAgent/session.py @@ -0,0 +1,842 @@ +"""Event-sourced sessions, persistence, replay, and context projection.""" + +from __future__ import annotations + +import json +import os +import sqlite3 +import threading +import hashlib +from copy import deepcopy +from dataclasses import asdict, dataclass, field, is_dataclass +from datetime import datetime, timezone +from pathlib import Path +from typing import Any, Callable, Iterable, Protocol +from uuid import uuid4 + + +SESSION_SCHEMA_VERSION = 1 +_SENSITIVE_KEYS = { + "api_key", + "apikey", + "authorization", + "credential", + "credentials", + "password", + "secret", + "token", +} + + +def _utc_now() -> str: + return datetime.now(timezone.utc).isoformat() + + +class SessionMigrationRegistry: + """Ordered migrations for persisted Session payloads and events.""" + + def __init__(self): + self._session_migrations: dict[int, Callable[[dict[str, Any]], dict[str, Any]]] = {} + self._event_migrations: dict[int, Callable[[dict[str, Any]], dict[str, Any]]] = {} + + def register_session(self, from_version: int, migration: Callable[[dict[str, Any]], dict[str, Any]]) -> None: + self._session_migrations[from_version] = migration + + def register_event(self, from_version: int, migration: Callable[[dict[str, Any]], dict[str, Any]]) -> None: + self._event_migrations[from_version] = migration + + def migrate(self, value: dict[str, Any]) -> dict[str, Any]: + payload = deepcopy(value) + version = int(payload.get("schema_version", 1)) + if version > SESSION_SCHEMA_VERSION: + raise ValueError( + f"unsupported Session schema_version={version}; maximum supported={SESSION_SCHEMA_VERSION}" + ) + while version < SESSION_SCHEMA_VERSION: + migration = self._session_migrations.get(version) + if migration is None: + raise ValueError(f"missing Session migration from schema_version={version}") + payload = migration(payload) + version = int(payload.get("schema_version", version + 1)) + payload["events"] = [self.migrate_event(event) for event in payload.get("events", [])] + return payload + + def migrate_event(self, value: dict[str, Any]) -> dict[str, Any]: + payload = deepcopy(value) + version = int(payload.get("schema_version", 1)) + if version > SESSION_SCHEMA_VERSION: + raise ValueError( + f"unsupported SessionEvent schema_version={version}; maximum supported={SESSION_SCHEMA_VERSION}" + ) + while version < SESSION_SCHEMA_VERSION: + migration = self._event_migrations.get(version) + if migration is None: + raise ValueError(f"missing SessionEvent migration from schema_version={version}") + payload = migration(payload) + version = int(payload.get("schema_version", version + 1)) + return payload + + +session_migrations = SessionMigrationRegistry() + + +def _json_safe(value: Any, *, key: str | None = None) -> Any: + """Return a JSON-safe copy while removing common credential fields.""" + if key is not None and key.lower() in _SENSITIVE_KEYS: + return "[redacted]" + if value is None or isinstance(value, (str, int, float, bool)): + return value + if is_dataclass(value): + return _json_safe(asdict(value)) + if hasattr(value, "to_dict") and callable(value.to_dict): + try: + return _json_safe(value.to_dict()) + except TypeError: + pass + if isinstance(value, dict): + return {str(item_key): _json_safe(item, key=str(item_key)) for item_key, item in value.items()} + if isinstance(value, (list, tuple, set)): + return [_json_safe(item) for item in value] + return repr(value) + + +@dataclass +class SessionEvent: + """One immutable fact in a session event log.""" + + type: str + session_id: str + data: dict[str, Any] = field(default_factory=dict) + event_id: str = field(default_factory=lambda: uuid4().hex) + timestamp: str = field(default_factory=_utc_now) + sequence: int = 0 + schema_version: int = SESSION_SCHEMA_VERSION + turn_id: str | None = None + step_id: str | None = None + run_id: str | None = None + agent_id: str | None = None + parent_event_id: str | None = None + + def __post_init__(self) -> None: + if not self.type or not isinstance(self.type, str): + raise ValueError("SessionEvent.type must be a non-empty string") + if not self.session_id: + raise ValueError("SessionEvent.session_id must not be empty") + if self.sequence < 0: + raise ValueError("SessionEvent.sequence must be non-negative") + if self.schema_version < 1: + raise ValueError("SessionEvent.schema_version must be at least 1") + self.data = _json_safe(self.data) + + def to_dict(self) -> dict[str, Any]: + return { + "event_id": self.event_id, + "session_id": self.session_id, + "type": self.type, + "data": deepcopy(self.data), + "timestamp": self.timestamp, + "sequence": self.sequence, + "schema_version": self.schema_version, + "turn_id": self.turn_id, + "step_id": self.step_id, + "run_id": self.run_id, + "agent_id": self.agent_id, + "parent_event_id": self.parent_event_id, + } + + @classmethod + def from_dict(cls, value: dict[str, Any]) -> "SessionEvent": + if not isinstance(value, dict): + raise TypeError("SessionEvent payload must be a dictionary") + return cls( + event_id=str(value.get("event_id") or uuid4().hex), + session_id=str(value["session_id"]), + type=str(value["type"]), + data=dict(value.get("data") or {}), + timestamp=str(value.get("timestamp") or _utc_now()), + sequence=int(value.get("sequence", 0)), + schema_version=int(value.get("schema_version", SESSION_SCHEMA_VERSION)), + turn_id=value.get("turn_id"), + step_id=value.get("step_id"), + run_id=value.get("run_id"), + agent_id=value.get("agent_id"), + parent_event_id=value.get("parent_event_id"), + ) + + +@dataclass +class SessionCheckpoint: + checkpoint_id: str + session_id: str + sequence: int + label: str | None = None + created_at: str = field(default_factory=_utc_now) + metadata: dict[str, Any] = field(default_factory=dict) + + def to_dict(self) -> dict[str, Any]: + return _json_safe(asdict(self)) + + +@dataclass +class SessionReplay: + session_id: str + event_count: int + completed_turns: list[str] = field(default_factory=list) + incomplete_turns: list[str] = field(default_factory=list) + failed_turns: list[str] = field(default_factory=list) + last_sequence: int = 0 + last_event_type: str | None = None + + def to_dict(self) -> dict[str, Any]: + return asdict(self) + + +@dataclass +class Session: + """Append-only session state reconstructed from versioned events.""" + + session_id: str = field(default_factory=lambda: uuid4().hex) + metadata: dict[str, Any] = field(default_factory=dict) + events: list[SessionEvent] = field(default_factory=list) + created_at: str = field(default_factory=_utc_now) + updated_at: str = field(default_factory=_utc_now) + schema_version: int = SESSION_SCHEMA_VERSION + + def __post_init__(self) -> None: + if not self.session_id: + raise ValueError("Session.session_id must not be empty") + self.metadata = _json_safe(self.metadata) + self.events = sorted(list(self.events), key=lambda event: event.sequence) + self.validate() + + @property + def next_sequence(self) -> int: + return self.events[-1].sequence + 1 if self.events else 1 + + def append( + self, + event_type: str, + data: dict[str, Any] | None = None, + *, + turn_id: str | None = None, + step_id: str | None = None, + run_id: str | None = None, + agent_id: str | None = None, + parent_event_id: str | None = None, + ) -> SessionEvent: + event = SessionEvent( + type=event_type, + session_id=self.session_id, + data=data or {}, + sequence=self.next_sequence, + schema_version=self.schema_version, + turn_id=turn_id, + step_id=step_id, + run_id=run_id, + agent_id=agent_id, + parent_event_id=parent_event_id, + ) + self.events.append(event) + self.updated_at = event.timestamp + return event + + def append_event(self, event: SessionEvent) -> SessionEvent: + if event.session_id != self.session_id: + raise ValueError("event session_id does not match Session") + expected = self.next_sequence + if event.sequence == 0: + event.sequence = expected + if event.sequence != expected: + raise ValueError(f"event sequence must be {expected}, got {event.sequence}") + if any(existing.event_id == event.event_id for existing in self.events): + raise ValueError(f"duplicate event_id: {event.event_id}") + self.events.append(event) + self.updated_at = event.timestamp + return event + + def page(self, *, after: int = 0, limit: int = 100) -> list[SessionEvent]: + if limit < 1: + raise ValueError("limit must be at least 1") + return [event for event in self.events if event.sequence > after][:limit] + + def checkpoint(self, label: str | None = None, metadata: dict[str, Any] | None = None) -> SessionCheckpoint: + checkpoint = SessionCheckpoint( + checkpoint_id=uuid4().hex, + session_id=self.session_id, + sequence=self.events[-1].sequence if self.events else 0, + label=label, + metadata=_json_safe(metadata or {}), + ) + self.append("session.checkpointed", {"checkpoint": checkpoint.to_dict()}) + return checkpoint + + def fork( + self, + *, + new_session_id: str | None = None, + through_sequence: int | None = None, + metadata: dict[str, Any] | None = None, + ) -> "Session": + boundary = through_sequence if through_sequence is not None else (self.events[-1].sequence if self.events else 0) + if boundary < 0 or boundary > (self.events[-1].sequence if self.events else 0): + raise ValueError("through_sequence is outside the session event range") + forked = Session( + session_id=new_session_id or uuid4().hex, + metadata={ + **deepcopy(self.metadata), + **(metadata or {}), + "forked_from": self.session_id, + "forked_at_sequence": boundary, + }, + ) + forked.append("session.started", {"forked_from": self.session_id, "through_sequence": boundary}) + for event in self.events: + if event.sequence > boundary: + break + forked.append( + event.type, + { + **deepcopy(event.data), + "source_session_id": self.session_id, + "source_event_id": event.event_id, + "source_sequence": event.sequence, + }, + turn_id=event.turn_id, + step_id=event.step_id, + run_id=event.run_id, + agent_id=event.agent_id, + parent_event_id=event.event_id, + ) + forked.append("session.forked", {"source_session_id": self.session_id, "through_sequence": boundary}) + return forked + + def replay(self) -> SessionReplay: + started: set[str] = set() + completed: set[str] = set() + failed: set[str] = set() + for event in self.events: + if not event.turn_id: + continue + if event.type == "turn.started": + started.add(event.turn_id) + elif event.type == "turn.completed": + completed.add(event.turn_id) + elif event.type == "turn.failed": + failed.add(event.turn_id) + return SessionReplay( + session_id=self.session_id, + event_count=len(self.events), + completed_turns=sorted(completed), + incomplete_turns=sorted(started - completed - failed), + failed_turns=sorted(failed), + last_sequence=self.events[-1].sequence if self.events else 0, + last_event_type=self.events[-1].type if self.events else None, + ) + + def validate(self) -> None: + seen_ids: set[str] = set() + expected = 1 + for event in self.events: + if event.session_id != self.session_id: + raise ValueError("session contains an event for another session") + if event.event_id in seen_ids: + raise ValueError(f"duplicate event_id: {event.event_id}") + if event.sequence != expected: + raise ValueError(f"event sequence gap: expected {expected}, got {event.sequence}") + if event.schema_version > SESSION_SCHEMA_VERSION: + raise ValueError( + f"unsupported SessionEvent schema_version={event.schema_version}; " + f"maximum supported={SESSION_SCHEMA_VERSION}" + ) + seen_ids.add(event.event_id) + expected += 1 + + def to_dict(self) -> dict[str, Any]: + return { + "session_id": self.session_id, + "metadata": deepcopy(self.metadata), + "events": [event.to_dict() for event in self.events], + "created_at": self.created_at, + "updated_at": self.updated_at, + "schema_version": self.schema_version, + } + + @classmethod + def from_dict(cls, value: dict[str, Any]) -> "Session": + value = session_migrations.migrate(value) + return cls( + session_id=str(value["session_id"]), + metadata=dict(value.get("metadata") or {}), + events=[SessionEvent.from_dict(event) for event in value.get("events", [])], + created_at=str(value.get("created_at") or _utc_now()), + updated_at=str(value.get("updated_at") or _utc_now()), + schema_version=int(value.get("schema_version", SESSION_SCHEMA_VERSION)), + ) + + +class SessionStore(Protocol): + def create(self, session: Session | None = None, *, metadata: dict[str, Any] | None = None) -> Session: + ... + + def get(self, session_id: str) -> Session | None: + ... + + def save(self, session: Session) -> None: + ... + + def list(self, *, limit: int = 100) -> list[Session]: + ... + + def delete(self, session_id: str) -> bool: + ... + + +class InMemorySessionStore: + """Thread-safe dependency-free SessionStore.""" + + def __init__(self): + self._sessions: dict[str, Session] = {} + self._lock = threading.RLock() + + def create(self, session: Session | None = None, *, metadata: dict[str, Any] | None = None) -> Session: + with self._lock: + value = session or Session(metadata=metadata or {}) + if value.session_id in self._sessions: + raise ValueError(f"session already exists: {value.session_id}") + self._sessions[value.session_id] = deepcopy(value) + return deepcopy(value) + + def get(self, session_id: str) -> Session | None: + with self._lock: + value = self._sessions.get(str(session_id)) + return deepcopy(value) if value is not None else None + + def save(self, session: Session) -> None: + session.validate() + with self._lock: + current = self._sessions.get(session.session_id) + if current and current.events and session.events: + if session.events[-1].sequence < current.events[-1].sequence: + raise ValueError("refusing to replace a Session with an older event sequence") + self._sessions[session.session_id] = deepcopy(session) + + def list(self, *, limit: int = 100) -> list[Session]: + if limit < 1: + raise ValueError("limit must be at least 1") + with self._lock: + values = sorted(self._sessions.values(), key=lambda item: item.updated_at, reverse=True) + return deepcopy(values[:limit]) + + def delete(self, session_id: str) -> bool: + with self._lock: + return self._sessions.pop(str(session_id), None) is not None + + +class JsonlSessionStore: + """One append-readable JSONL file per Session with atomic full saves.""" + + def __init__(self, directory: str | os.PathLike[str]): + self.directory = Path(directory) + self.directory.mkdir(parents=True, exist_ok=True) + self._lock = threading.RLock() + + @staticmethod + def _safe_id(session_id: str) -> str: + value = str(session_id) + if not value or any(character not in "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789-_" for character in value): + raise ValueError("session_id contains unsafe characters") + return value + + def _path(self, session_id: str) -> Path: + return self.directory / f"{self._safe_id(session_id)}.jsonl" + + def create(self, session: Session | None = None, *, metadata: dict[str, Any] | None = None) -> Session: + value = session or Session(metadata=metadata or {}) + with self._lock: + if self._path(value.session_id).exists(): + raise ValueError(f"session already exists: {value.session_id}") + self.save(value) + return deepcopy(value) + + def get(self, session_id: str) -> Session | None: + path = self._path(session_id) + if not path.exists(): + return None + with self._lock, path.open("r", encoding="utf-8") as handle: + lines = [json.loads(line) for line in handle if line.strip()] + if not lines or lines[0].get("kind") != "session": + raise ValueError(f"invalid Session JSONL header: {path}") + header = lines[0]["value"] + events = [SessionEvent.from_dict(line["value"]) for line in lines[1:] if line.get("kind") == "event"] + return Session( + session_id=header["session_id"], + metadata=header.get("metadata") or {}, + events=events, + created_at=header.get("created_at") or _utc_now(), + updated_at=header.get("updated_at") or _utc_now(), + schema_version=int(header.get("schema_version", SESSION_SCHEMA_VERSION)), + ) + + def save(self, session: Session) -> None: + session.validate() + path = self._path(session.session_id) + temporary = path.with_suffix(f".tmp-{uuid4().hex}") + header = { + "session_id": session.session_id, + "metadata": session.metadata, + "created_at": session.created_at, + "updated_at": session.updated_at, + "schema_version": session.schema_version, + } + with self._lock: + with temporary.open("w", encoding="utf-8") as handle: + handle.write(json.dumps({"kind": "session", "value": header}, ensure_ascii=False) + "\n") + for event in session.events: + handle.write(json.dumps({"kind": "event", "value": event.to_dict()}, ensure_ascii=False) + "\n") + handle.flush() + os.fsync(handle.fileno()) + os.replace(temporary, path) + + def list(self, *, limit: int = 100) -> list[Session]: + if limit < 1: + raise ValueError("limit must be at least 1") + sessions = [] + for path in sorted(self.directory.glob("*.jsonl"), key=lambda item: item.stat().st_mtime, reverse=True): + session = self.get(path.stem) + if session is not None: + sessions.append(session) + if len(sessions) >= limit: + break + return sessions + + def delete(self, session_id: str) -> bool: + path = self._path(session_id) + with self._lock: + if not path.exists(): + return False + path.unlink() + return True + + +class SqliteSessionStore: + """Standard-library SQLite SessionStore with transactional saves.""" + + def __init__(self, path: str | os.PathLike[str]): + self.path = str(path) + Path(self.path).parent.mkdir(parents=True, exist_ok=True) + self._lock = threading.RLock() + self._initialize() + + def _connect(self) -> sqlite3.Connection: + connection = sqlite3.connect(self.path, timeout=30) + connection.row_factory = sqlite3.Row + return connection + + def _initialize(self) -> None: + with self._connect() as connection: + connection.execute("PRAGMA journal_mode=WAL") + connection.execute( + """ + CREATE TABLE IF NOT EXISTS sessions ( + session_id TEXT PRIMARY KEY, + metadata TEXT NOT NULL, + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL, + schema_version INTEGER NOT NULL + ) + """ + ) + connection.execute( + """ + CREATE TABLE IF NOT EXISTS session_events ( + session_id TEXT NOT NULL, + sequence INTEGER NOT NULL, + event_id TEXT NOT NULL UNIQUE, + payload TEXT NOT NULL, + PRIMARY KEY (session_id, sequence), + FOREIGN KEY (session_id) REFERENCES sessions(session_id) ON DELETE CASCADE + ) + """ + ) + + def create(self, session: Session | None = None, *, metadata: dict[str, Any] | None = None) -> Session: + value = session or Session(metadata=metadata or {}) + with self._lock, self._connect() as connection: + existing = connection.execute( + "SELECT 1 FROM sessions WHERE session_id = ?", (value.session_id,) + ).fetchone() + if existing: + raise ValueError(f"session already exists: {value.session_id}") + self.save(value) + return deepcopy(value) + + def get(self, session_id: str) -> Session | None: + with self._lock, self._connect() as connection: + row = connection.execute( + "SELECT * FROM sessions WHERE session_id = ?", (str(session_id),) + ).fetchone() + if row is None: + return None + event_rows = connection.execute( + "SELECT payload FROM session_events WHERE session_id = ? ORDER BY sequence", + (str(session_id),), + ).fetchall() + return Session( + session_id=row["session_id"], + metadata=json.loads(row["metadata"]), + events=[SessionEvent.from_dict(json.loads(event_row["payload"])) for event_row in event_rows], + created_at=row["created_at"], + updated_at=row["updated_at"], + schema_version=row["schema_version"], + ) + + def save(self, session: Session) -> None: + session.validate() + with self._lock, self._connect() as connection: + connection.execute("PRAGMA foreign_keys=ON") + connection.execute( + """ + INSERT INTO sessions(session_id, metadata, created_at, updated_at, schema_version) + VALUES (?, ?, ?, ?, ?) + ON CONFLICT(session_id) DO UPDATE SET + metadata=excluded.metadata, + updated_at=excluded.updated_at, + schema_version=excluded.schema_version + """, + ( + session.session_id, + json.dumps(session.metadata, ensure_ascii=False), + session.created_at, + session.updated_at, + session.schema_version, + ), + ) + connection.execute("DELETE FROM session_events WHERE session_id = ?", (session.session_id,)) + connection.executemany( + "INSERT INTO session_events(session_id, sequence, event_id, payload) VALUES (?, ?, ?, ?)", + [ + ( + session.session_id, + event.sequence, + event.event_id, + json.dumps(event.to_dict(), ensure_ascii=False), + ) + for event in session.events + ], + ) + + def list(self, *, limit: int = 100) -> list[Session]: + if limit < 1: + raise ValueError("limit must be at least 1") + with self._lock, self._connect() as connection: + rows = connection.execute( + "SELECT session_id FROM sessions ORDER BY updated_at DESC LIMIT ?", (int(limit),) + ).fetchall() + return [session for row in rows if (session := self.get(row["session_id"])) is not None] + + def delete(self, session_id: str) -> bool: + with self._lock, self._connect() as connection: + connection.execute("PRAGMA foreign_keys=ON") + cursor = connection.execute("DELETE FROM sessions WHERE session_id = ?", (str(session_id),)) + return cursor.rowcount > 0 + + +class ContextProjector: + """Project model messages and runtime views from Session events.""" + + def messages(self, session: Session, *, turn_id: str | None = None) -> list[dict[str, Any]]: + requested = [ + event for event in session.events + if event.type == "model.requested" + and (turn_id is None or event.turn_id == turn_id) + and isinstance(event.data.get("messages"), list) + ] + messages: list[dict[str, Any]] = [] + for event in session.events: + if turn_id is not None and event.turn_id != turn_id: + continue + if event.type in {"message.received", "message.created"}: + messages.append({"role": event.data.get("role", "user"), "content": event.data.get("content", "")}) + elif event.type in {"assistant.completed", "model.completed"} and event.data.get("content") is not None: + messages.append({"role": "assistant", "content": event.data.get("content", "")}) + elif event.type == "tool.completed": + messages.append({ + "role": "tool", + "tool_call_id": event.data.get("tool_call_id"), + "content": str(event.data.get("output", "")), + }) + if messages: + return messages + if requested: + return deepcopy(requested[-1].data["messages"]) + return [] + + def model_request( + self, + session: Session, + *, + sequence: int | None = None, + ) -> list[dict[str, Any]]: + """Reconstruct the exact messages sent for one persisted model request.""" + requested = [ + event for event in session.events + if event.type == "model.requested" + and isinstance(event.data.get("messages"), list) + and (sequence is None or event.sequence == sequence) + ] + if not requested: + raise LookupError("model request event was not found") + return deepcopy(requested[-1].data["messages"]) + + def trace(self, session: Session, *, turn_id: str | None = None) -> list[dict[str, Any]]: + trace = [] + for event in session.events: + if turn_id is not None and event.turn_id != turn_id: + continue + trace_type = event.data.get("trace_type") + if trace_type: + trace.append({ + "type": trace_type, + "data": deepcopy(event.data.get("trace_data") or {}), + "timestamp": event.timestamp, + "trace_id": event.data.get("trace_id"), + "parent_trace_id": event.data.get("parent_trace_id"), + "run_group_id": event.data.get("run_group_id"), + }) + return trace + + +@dataclass +class CompactionResult: + messages: list[dict[str, Any]] + removed_count: int + summary: str | None = None + spilled: list[dict[str, Any]] = field(default_factory=list) + + +@dataclass(frozen=True) +class ContextBudget: + """Dependency-free context budget using a conservative character estimate.""" + + max_tokens: int = 8192 + reserved_output_tokens: int = 1024 + chars_per_token: float = 4.0 + + def __post_init__(self) -> None: + if self.max_tokens < 1 or self.reserved_output_tokens < 0: + raise ValueError("context token limits must be non-negative") + if self.reserved_output_tokens >= self.max_tokens: + raise ValueError("reserved_output_tokens must be smaller than max_tokens") + if self.chars_per_token <= 0: + raise ValueError("chars_per_token must be positive") + + def estimate(self, messages: Iterable[dict[str, Any]]) -> int: + rendered = json.dumps(list(messages), ensure_ascii=False, sort_keys=True) + return max(1, int(len(rendered) / self.chars_per_token) + 1) + + def fits(self, messages: Iterable[dict[str, Any]]) -> bool: + return self.estimate(messages) <= self.max_tokens - self.reserved_output_tokens + + +class ContextCompactor: + """Deterministic message trimming with optional summary generation.""" + + def __init__( + self, + summarizer: Callable[[list[dict[str, Any]]], str] | None = None, + *, + max_inline_tool_chars: int = 12000, + ): + self.summarizer = summarizer + self.max_inline_tool_chars = max_inline_tool_chars + + def compact(self, messages: Iterable[dict[str, Any]], *, max_messages: int = 20) -> CompactionResult: + values, spilled = self._spill_tool_results(messages) + if max_messages < 2: + raise ValueError("max_messages must be at least 2") + if len(values) <= max_messages: + return CompactionResult(messages=values, removed_count=0, spilled=spilled) + + system = [message for message in values if message.get("role") == "system"][:1] + non_system = [message for message in values if message.get("role") != "system"] + keep_count = max_messages - len(system) + kept = non_system[-keep_count:] + removed = non_system[:-keep_count] + + while kept and kept[0].get("role") == "tool": + removed.append(kept.pop(0)) + summary = self.summarizer(deepcopy(removed)) if self.summarizer and removed else None + if summary: + summary_message = { + "role": "system", + "content": f"Previous conversation summary:\n{summary}", + "metadata": {"compacted": True, "removed_count": len(removed)}, + } + result_messages = system + [summary_message] + kept + else: + result_messages = system + kept + return CompactionResult( + messages=result_messages, + removed_count=len(removed), + summary=summary, + spilled=spilled, + ) + + def compact_to_budget( + self, + messages: Iterable[dict[str, Any]], + budget: ContextBudget, + ) -> CompactionResult: + values = list(messages) + if budget.fits(values): + prepared, spilled = self._spill_tool_results(values) + return CompactionResult(messages=prepared, removed_count=0, spilled=spilled) + max_messages = max(2, len(values)) + while max_messages > 2: + result = self.compact(values, max_messages=max_messages) + if budget.fits(result.messages): + return result + max_messages -= 1 + return self.compact(values, max_messages=2) + + def _spill_tool_results( + self, + messages: Iterable[dict[str, Any]], + ) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]: + prepared: list[dict[str, Any]] = [] + spilled: list[dict[str, Any]] = [] + for message in messages: + value = deepcopy(message) + content = value.get("content") + if value.get("role") == "tool" and isinstance(content, str) and len(content) > self.max_inline_tool_chars: + digest = hashlib.sha256(content.encode("utf-8")).hexdigest() + reference = f"tool-result:sha256:{digest}" + spilled.append({ + "reference": reference, + "sha256": digest, + "size": len(content), + "content": content, + }) + value["content"] = f"[Tool result stored outside context: {reference}, {len(content)} chars]" + value.setdefault("metadata", {})["content_reference"] = reference + prepared.append(value) + return prepared, spilled + + +__all__ = [ + "SESSION_SCHEMA_VERSION", + "SessionMigrationRegistry", + "session_migrations", + "SessionEvent", + "SessionCheckpoint", + "SessionReplay", + "Session", + "SessionStore", + "InMemorySessionStore", + "JsonlSessionStore", + "SqliteSessionStore", + "ContextProjector", + "CompactionResult", + "ContextBudget", + "ContextCompactor", +] diff --git a/LightAgent/skills.py b/LightAgent/skills.py index a2bb81f..b6c767a 100644 --- a/LightAgent/skills.py +++ b/LightAgent/skills.py @@ -27,6 +27,8 @@ class Skill: has_scripts: bool = False has_references: bool = False has_assets: bool = False + source_directory: str | None = None + precedence: int = 0 class SkillManager: @@ -35,17 +37,20 @@ class SkillManager: def __init__(self, skills_directories: List[str] = None, logger=None): self.skills_directories = skills_directories or ["skills"] self.skills: Dict[str, Skill] = {} + self.skill_conflicts: List[Dict[str, Any]] = [] self.logger = logger or logging.getLogger(__name__) def discover_skills(self) -> List[Skill]: """发现所有可用技能(仅加载元数据)""" discovered = [] + self.skills = {} + self.skill_conflicts = [] - for base_dir in self.skills_directories: + for precedence, base_dir in enumerate(self.skills_directories): if not os.path.exists(base_dir): continue - for item in os.listdir(base_dir): + for item in sorted(os.listdir(base_dir)): skill_path = os.path.join(base_dir, item) skill_file = os.path.join(skill_path, "SKILL.md") @@ -53,6 +58,18 @@ def discover_skills(self) -> List[Skill]: try: skill = self._load_skill_metadata(skill_path) if skill: + skill.source_directory = str(Path(base_dir).resolve()) + skill.precedence = precedence + previous = self.skills.get(skill.name) + if previous is not None: + conflict = { + "name": skill.name, + "winner": skill.path, + "shadowed": previous.path, + "rule": "later skills_directories entry wins", + } + self.skill_conflicts.append(conflict) + self._log("WARNING", "skill_conflict", conflict) self.skills[skill.name] = skill discovered.append(skill) self._log("DEBUG", "discover_skill", @@ -63,6 +80,33 @@ def discover_skills(self) -> List[Skill]: return discovered + def list_conflicts(self) -> List[Dict[str, Any]]: + """Return deterministic diagnostics for shadowed Skill names.""" + return [dict(item) for item in self.skill_conflicts] + + def discover_project_instructions( + self, + start_directory: str | os.PathLike[str] | None = None, + *, + filename: str = "AGENTS.md", + max_chars: int = 100_000, + ) -> str: + """Load project instructions from filesystem root to the working directory.""" + current = Path(start_directory or os.getcwd()).resolve() + candidates = [current, *current.parents] + contents: List[str] = [] + total = 0 + for directory in reversed(candidates): + instruction_file = directory / filename + if not instruction_file.is_file(): + continue + content = instruction_file.read_text(encoding="utf-8") + total += len(content) + if total > max_chars: + raise ValueError(f"project instructions exceed max_chars={max_chars}") + contents.append(f"# {instruction_file}\n{content.strip()}") + return "\n\n".join(contents) + def _load_skill_metadata(self, skill_path: str) -> Optional[Skill]: """从SKILL.md加载技能元数据(仅frontmatter)""" skill_file = os.path.join(skill_path, "SKILL.md") @@ -136,20 +180,21 @@ def execute_script(self, skill_name: str, script_name: str, args: List[str] = No return f"Error: Skill '{skill_name}' not found" skill = self.skills[skill_name] - script_path = os.path.join(skill.path, "scripts", script_name) - - if not os.path.exists(script_path): - return f"Error: Script '{script_path}' not found in skill '{skill_name}'" + scripts_root = Path(skill.path, "scripts").resolve() + script_path = (scripts_root / script_name).resolve() # 安全检查:只允许执行scripts目录下的文件 - if not script_path.startswith(os.path.join(skill.path, "scripts")): + if not script_path.is_relative_to(scripts_root): return "Error: Security violation - cannot execute outside scripts directory" + if not script_path.exists(): + return f"Error: Script '{script_path}' not found in skill '{skill_name}'" + try: # 在临时目录中执行以提供隔离 with tempfile.TemporaryDirectory() as tmpdir: result = subprocess.run( - [script_path] + (args or []), + [str(script_path)] + (args or []), capture_output=True, text=True, timeout=30, @@ -175,13 +220,14 @@ def read_reference(self, skill_name: str, ref_path: str) -> str: return f"Error: Skill '{skill_name}' not found" skill = self.skills[skill_name] - full_path = os.path.join(skill.path, "references", ref_path) + references_root = Path(skill.path, "references").resolve() + full_path = (references_root / ref_path).resolve() # 安全检查:防止目录遍历 - if not full_path.startswith(os.path.join(skill.path, "references")): + if not full_path.is_relative_to(references_root): return "Error: Security violation - invalid reference path" - if not os.path.exists(full_path): + if not full_path.exists(): return f"Error: Reference '{ref_path}' not found" try: @@ -197,10 +243,11 @@ def read_asset(self, skill_name: str, asset_path: str) -> bytes: raise ValueError(f"Skill '{skill_name}' not found") skill = self.skills[skill_name] - full_path = os.path.join(skill.path, "assets", asset_path) + assets_root = Path(skill.path, "assets").resolve() + full_path = (assets_root / asset_path).resolve() # 安全检查 - if not full_path.startswith(os.path.join(skill.path, "assets")): + if not full_path.is_relative_to(assets_root): raise ValueError("Security violation: invalid asset path") with open(full_path, 'rb') as f: diff --git a/LightAgent/version.py b/LightAgent/version.py index 2ac1ae2..1e299c8 100644 --- a/LightAgent/version.py +++ b/LightAgent/version.py @@ -6,4 +6,4 @@ 最后更新: 2026-08-09 """ -__version__ = "0.9.7" +__version__ = "0.10.0" diff --git a/README.md b/README.md index 94aca76..c389366 100644 --- a/README.md +++ b/README.md @@ -37,7 +37,8 @@ LightAgent is an ultra‑lightweight, open‑source framework that now natively --- ## News -- new**[2026-08-15]** LightAgent v0.9.7 Released: Adds a dependency-free Connector contract with offline validation and examples, expands Python executor security checks, introduces an opt-in Mem0 Graph security matrix, and adds a public API compatibility inventory for v1.0 stabilization. +- new**[2026-08-15]** LightAgent v0.10.0 Development: Adds the unified event-sourced Agent Runtime with durable Sessions, async execution, Capability Registry and Policy, Inbox/Goals/Budgets, compaction and recovery, Jobs/subagents, standardized Skills/MCP adapters, and SQLite FTS5 retrieval. +- **[2026-08-15]** LightAgent v0.9.7 Released: Adds a dependency-free Connector contract with offline validation and examples, expands Python executor security checks, introduces an opt-in Mem0 Graph security matrix, and adds a public API compatibility inventory for v1.0 stabilization. - **[2026-07-30]** LightAgent v0.9.6 Released: Adds production trace summaries and exporters, deterministic evaluation, durable human approval for tools, handoffs, and LightFlow, plus fail-closed shared Graph Memory admission and audit controls. - **[2026-07-10]** LightAgent v0.9.3 Released: Completes runtime hook lifecycle coverage and hardens streaming tool safety with `max_tool_iterations`, consistent `on_error` / `after_run` closure, and expanded regression coverage. - **[2026-06-24]** LightAgent v0.9.0 Development: Adds checkpointed LightFlow workflows with resume/rerun support, approval nodes, richer step status and trace metadata, reusable Guardrails templates, stronger MemoryPolicy controls, and the first SharedMemoryPool prototype. @@ -70,6 +71,7 @@ Older release notes are available on [GitHub Releases](https://github.com/wanxin - **Evaluation Harness** 📊: `LightEvaluator` runs deterministic agent or LightFlow regression cases for output, tool choice, policy events, recovery, latency, usage, and estimated cost. - **Human Review** 👤: `HumanApprovalHook`, durable review stores, and LightFlow approval checkpoints support approve, reject, argument editing, human responses, batches, and trace feedback for high-impact actions. - **Runtime Hooks** 🧩: Ordered `hooks=[...]` middleware can observe, replace, or block run, model, tool, memory, and LightFlow step phases while recording hook decisions in trace events. +- **Event-Sourced Runtime** 🧱: Optional durable Sessions, replay, checkpoints, forks, async entry points, scoped Capability Providers, unified Policy, Inbox, Goals, Budgets, Jobs, subagents, context compaction, and SQLite FTS5 retrieval. - **Guardrails Templates** 🛡️: Reusable input/tool/output guardrail templates help block private data, require confirmation for sensitive tools, validate high-risk parameters, and redact sensitive output. - **Tool Generator** 🚀: Just provide your API documentation to the [Tool Generator], which will automatically create exclusive tools for you, allowing you to quickly build hundreds of personalized custom tools in just 1 hour to improve efficiency and unleash your creative potential. - **Agent Self-Learning** 🧠️: Each agent has its own scene memory capabilities and the ability to self-learn from user conversations. @@ -80,6 +82,8 @@ Older release notes are available on [GitHub Releases](https://github.com/wanxin | Layer | Main API | Use it when you need | | --- | --- | --- | | Single agent runtime | `LightAgent` | One agent with model calls, tools, memory, streaming, trace, and guardrails. | +| Durable runtime state | `Session`, `AgentRuntime` | Replayable events, checkpoints, Inbox, Goals, Budgets, Jobs, and context recovery. | +| Capability layer | `CapabilityRegistry`, `PolicyEngine` | Scoped Providers, permission snapshots, lifecycle, policy, and audit. | | Multi-agent routing | `LightSwarm` | Role-based delegation across specialized agents. | | Deterministic workflow | `LightFlow` | Ordered DAG workflows, retries, checkpoints, durable approvals, resume, and rerun. | | Tools and integrations | `tools`, `ToolRegistry`, MCP | Python tools, generated tools, runtime tool loading, or MCP tool servers. | @@ -108,6 +112,8 @@ LightAgent keeps the default call path simple while allowing production controls | Evaluation | `LightEvaluator().run(agent, cases)` | Run deterministic behavioral checks from structured traces. | | Tool approval | `LightAgent(..., hooks=[HumanApprovalHook(...)])` | Require review before selected tools or handoffs. | | Workflow | `LightFlow().step(...).run(query)` | Use for deterministic multi-step execution. | +| Durable session | `agent.run(query, session_id="project-42")` | Continue and replay a persisted conversation. | +| Async | `await agent.arun(query)` | Keep an asyncio application responsive. | ### Evaluate And Review High-Risk Actions @@ -175,6 +181,8 @@ For tool/handoff approval, durable LightFlow review, batches, and feedback, see For the v1.0 stability proposal, supported Python versions, public imports, and compatibility promises, see [Public API And Compatibility Inventory](docs/public_api_compatibility.md). +For durable Sessions, Capability Providers, Policy, Inbox, Goals, Budgets, Jobs, compaction, subagents, Skills/MCP updates, and SQLite FTS5 retrieval, see [LightAgent v0.10 Runtime](docs/runtime_v010.md). + For browser-use integration with recent `browser-use` versions, see [browser-use Integration](docs/browser_use.md). --- diff --git a/README.zh-CN.md b/README.zh-CN.md index 8dbc027..9aa8b29 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -55,7 +55,8 @@ --- ## 新闻 -- new**[2026-08-15]** LightAgent v0.9.7 正式发布:新增轻量、零依赖的 Connector 契约及离线校验与示例,强化 Python 执行器安全检查,补充可选的 Mem0 Graph 安全验证矩阵,并通过公共 API 兼容性清单为 v1.0 稳定化做准备。 +- new**[2026-08-15]** LightAgent v0.10.0 开发版:新增统一的事件溯源 Agent Runtime,支持可持久化 Session、异步执行、Capability Registry 与 Policy、Inbox/Goals/Budgets、上下文压缩与恢复、Jobs/子 Agent、标准化 Skills/MCP 适配器和 SQLite FTS5 检索。 +- **[2026-08-15]** LightAgent v0.9.7 正式发布:新增轻量、零依赖的 Connector 契约及离线校验与示例,强化 Python 执行器安全检查,补充可选的 Mem0 Graph 安全验证矩阵,并通过公共 API 兼容性清单为 v1.0 稳定化做准备。 - **[2026-07-30]** LightAgent v0.9.6 正式发布:新增生产级 Trace 汇总与导出、确定性评测、工具、handoff 和 LightFlow 的可持久化人工审批,以及共享 Graph Memory 的 fail-closed 写入准入与审计控制。 - **[2026-07-10]** LightAgent v0.9.3 正式发布:补全 `after_run`、`on_error` 和记忆检索 Hooks 生命周期,新增独立的 `max_tool_iterations` 流式工具循环上限并保持 `max_retry` 向后兼容,同时完善错误收尾、Trace 和回归测试。 - **[2026-06-24]** LightAgent v0.9.0 开发版:新增可持久化 LightFlow checkpoint、resume/rerun、审批节点、更清晰的步骤状态和 trace 元数据,同时补充 Guardrails 模板、MemoryPolicy 控制和 SharedMemoryPool 原型。 @@ -85,6 +86,7 @@ - **评测框架** 📊:`LightEvaluator` 可对 Agent 或 LightFlow 执行确定性回归用例,检查输出、工具选择、策略事件、恢复能力、延迟、usage 和预估成本。 - **人工审核** 👤:`HumanApprovalHook`、可持久化审核存储和 LightFlow 审批 checkpoint 支持高风险动作的批准、拒绝、参数编辑、人工响应、批量审核和 trace 反馈。 - **Runtime Hooks** 🧩:通过有序 `hooks=[...]` 中间件观察、替换或阻断运行、模型、工具、记忆、LightSwarm handoff 和 LightFlow 步骤阶段;安全策略可使用 `PolicyHook` 在异常或超时时 fail closed,并将决策写入 trace。 +- **事件溯源 Runtime** 🧱:可选的持久化 Session、回放、checkpoint、fork、异步入口、分层 Capability Provider、统一 Policy、Inbox、Goals、Budgets、Jobs、子 Agent、上下文压缩和 SQLite FTS5 检索。 - **Tools工具生成器** 🚀:只需将您的API文档交给[[Tools工具生成器]](#3-tools工具生成器),它将自动化地为您打造专属的tools,助您在短短1小时内快速构建数百个个性化的自定义工具,提升效率,释放您的创新潜能。 - **agent自我学习** 🧠️:每个agent拥有自己的场景记忆能力,拥有从用户的对话中进行自我学习能力。 - **自适应tools机制** 🛠️:支持添加无限量tools,在上万个工具中让大模型过滤无关工具后再发送给大模型,可大幅度降低Token消耗。 @@ -97,6 +99,8 @@ | 层级 | 主要 API | 适用场景 | | --- | --- | --- | | 单 Agent 运行时 | `LightAgent` | 一个 Agent 的模型调用、工具、记忆、流式输出、trace 和 guardrails。 | +| 持久化运行状态 | `Session`、`AgentRuntime` | 可回放事件、checkpoint、Inbox、Goals、Budgets、Jobs 和上下文恢复。 | +| 能力层 | `CapabilityRegistry`、`PolicyEngine` | 分层 Provider、权限快照、生命周期、策略与审计。 | | 多 Agent 路由 | `LightSwarm` | 在多个专业 Agent 之间进行角色化委托。 | | 确定性工作流 | `LightFlow` | DAG 工作流、重试、checkpoint、持久化审批、resume 和 rerun。 | | 工具与集成 | `tools`、`ToolRegistry`、MCP | Python 工具、生成工具、运行时加载工具或 MCP 工具服务。 | @@ -125,6 +129,8 @@ LightAgent 保持默认调用路径简单,同时允许逐步加入生产级控 | 评测 | `LightEvaluator().run(agent, cases)` | 基于结构化 trace 执行确定性行为检查。 | | 工具审批 | `LightAgent(..., hooks=[HumanApprovalHook(...)])` | 在选定工具或 handoff 执行前要求人工审核。 | | 工作流 | `LightFlow().step(...).run(query)` | 用于确定性多步骤执行。 | +| 持久化 Session | `agent.run(query, session_id="project-42")` | 延续并回放持久化会话。 | +| 异步调用 | `await agent.arun(query)` | 避免阻塞 asyncio 应用。 | ### 评测并审核高风险动作 @@ -162,6 +168,7 @@ print(report.to_dict()) - Trace 可观测能力请查看 [Trace Observability](docs/tracing.md)。 - 确定性回归用例、指标和 CI 方案请查看 [Evaluation Harness](docs/evaluation.md)。 - 工具/handoff 审批、LightFlow 持久化审核、批量决策和反馈请查看 [Human Review](docs/human_review.md)。 +- 持久化 Session、Capability Provider、Policy、Inbox、Goals、Budgets、Jobs、上下文压缩、子 Agent、Skills/MCP 更新和 SQLite FTS5 检索,请查看 [LightAgent v0.10 Runtime](docs/runtime_v010.md)。 --- diff --git a/docs/runtime_v010.md b/docs/runtime_v010.md new file mode 100644 index 0000000..1ab5654 --- /dev/null +++ b/docs/runtime_v010.md @@ -0,0 +1,140 @@ +# LightAgent v0.10 Runtime + +LightAgent v0.10 adds an opt-in, event-sourced runtime beneath the compatible +`agent.run()` API. Basic agents still require no database or runtime setup. + +## Durable Sessions + +```python +from LightAgent import LightAgent, SqliteSessionStore + +store = SqliteSessionStore(".lightagent/sessions.sqlite3") +agent = LightAgent( + model="deepseek-v4-flash", + api_key="your_api_key", + base_url="your_base_url", + session_store=store, +) + +agent.run("Remember this project decision", session_id="project-42") +agent.run("What did we decide?", session_id="project-42") + +print(agent.replay_session("project-42")) +checkpoint = agent.checkpoint_session("before-implementation") +fork = agent.fork_session(through_sequence=checkpoint["sequence"]) +``` + +`InMemorySessionStore`, `JsonlSessionStore`, and `SqliteSessionStore` implement +the same contract. Events are append-only, sequence-validated, versioned, and +credential fields are redacted before persistence. `ContextProjector` rebuilds +conversation context, exact model requests, and compatible trace views from +the event log. + +## Async Usage + +```python +result = await agent.arun("Analyze the incident") + +stream = await agent.arun("Analyze the incident", stream=True) +async for chunk in stream: + print(chunk) +``` + +The async entry point keeps synchronous v0.9 model and tool clients compatible +by isolating them in a worker thread. Native async Providers and background +Jobs execute on the caller's event loop. + +## Capabilities And Policy + +```python +from LightAgent import ( + BaseCapabilityProvider, + CapabilityRegistry, + CapabilitySpec, + PolicyDecision, + PolicyEngine, +) + +def workspace_policy(request): + if request.capability.write and request.context.metadata.get("read_only"): + return PolicyDecision.block("workspace is read-only") + return PolicyDecision.allow(request.arguments) + +registry = CapabilityRegistry(policy_engine=PolicyEngine([workspace_policy])) +``` + +Providers declare lifecycle methods and capability metadata for read, write, +network, execution, persistence, risk, timeout, output limits, cancellation, +and approval. Runtime, Session, and Agent scopes resolve deterministically; +equal-scope conflicts are available through `registry.conflicts()`. + +`ToolProviderAdapter`, `MemoryProviderAdapter`, `SkillProviderAdapter`, +`MCPProviderAdapter`, and `WorkflowProviderAdapter` bridge existing LightAgent +APIs into the registry. Sensitive tool calls made by `LightAgent` pass through +the registry's `PolicyEngine` before existing Guardrails and Hooks. + +## Long-Task State + +Every `LightAgent` exposes `agent.runtime`: + +```python +goal = agent.runtime.goals.create( + "Prepare the release", + acceptance_criteria=["tests pass", "artifacts build"], +) +agent.runtime.goals.activate(goal.goal_id) +agent.runtime.inbox.enqueue("steering", "Do not publish yet", message_id="release-hold") +agent.runtime.pause("waiting for CI") +``` + +The runtime includes ordered and idempotent Inbox messages, durable Goals, +model/tool/token/time/cost budgets, progress-loop detection, cancellable Jobs, +and bounded subagent registration with narrowing-only permission snapshots. +Restoring a Session rebuilds Inbox, Goal, Budget, and interrupted Job state. + +## Context And Knowledge + +`ContextBudget` and `ContextCompactor` provide deterministic trimming, +optional summarization, and SHA-256 references for oversized tool output. +Compaction decisions are Session events. Checkpoints and forks retain source +event lineage. + +`SqliteFTSRetrievalProvider` is a dependency-free minimum RAG implementation: + +```python +from LightAgent import RetrievalDocument, SqliteFTSRetrievalProvider + +rag = SqliteFTSRetrievalProvider(".lightagent/knowledge.sqlite3") +rag.ingest(RetrievalDocument( + content="The production deployment uses blue-green releases.", + title="Deployment guide", + source="docs/deployment.md", + tenant_id="acme", +)) + +for result in rag.search("deployment", tenant_id="acme"): + print(result.citation_id, result.content) +``` + +`SessionSearchProvider` is intentionally separate from knowledge-base RAG and +returns citations in `session::` form. + +## Skills And MCP + +Skill directories are processed in declared order; later directories shadow +earlier entries and conflicts are reported by `SkillManager.list_conflicts()`. +`discover_project_instructions()` loads nested `AGENTS.md` files from root to +the working directory without modifying them. + +MCP retains stdio and SSE configuration compatibility and adds opt-in +Streamable HTTP (`transport: streamable-http`), reconnect attempts, tool-list +refresh, namespace isolation, and external credential-header resolution. + +## Compatibility Boundary + +- `agent.run("hello")`, `stream=True`, structured results, Hooks, Guardrails, + Memory, LightSwarm, and LightFlow remain compatible. +- Browser, Terminal, Shell, LSP, vector database, hosted service, and WebUI + implementations remain optional Providers rather than core dependencies. +- Session and Capability APIs are pre-1.0 contracts and may receive additive + changes before the v1.0 API freeze. diff --git a/example/13.runtime_v010.py b/example/13.runtime_v010.py new file mode 100644 index 0000000..0df4e17 --- /dev/null +++ b/example/13.runtime_v010.py @@ -0,0 +1,38 @@ +#!/usr/bin/env python +"""Minimal v0.10 Session, Policy, and SQLite RAG example.""" + +import os +import tempfile +from pathlib import Path + +from LightAgent import ( + BudgetLimits, + JsonlSessionStore, + LightAgent, + RetrievalDocument, + SqliteFTSRetrievalProvider, +) + + +with tempfile.TemporaryDirectory() as directory: + root = Path(directory) + agent = LightAgent( + model=os.getenv("LIGHTAGENT_MODEL", "deepseek-v4-flash"), + api_key=os.getenv("LIGHTAGENT_API_KEY", "your_api_key"), + base_url=os.getenv("LIGHTAGENT_BASE_URL", "your_base_url"), + session_store=JsonlSessionStore(root / "sessions"), + budget_limits=BudgetLimits(model_calls=10, tool_calls=20), + auto_discover_skills=False, + ) + + # A real call persists a complete Turn under this stable Session ID: + # print(agent.run("Summarize today's work", session_id="demo-session")) + + rag = SqliteFTSRetrievalProvider(root / "knowledge.sqlite3") + rag.ingest(RetrievalDocument( + title="Release policy", + source="docs/release.md", + content="Every release requires a passing full test suite.", + )) + for result in rag.search("release"): + print(result.citation_id, result.content) diff --git a/pyproject.toml b/pyproject.toml index 34dfa6f..0e08714 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [tool.poetry] name = "LightAgent" -version = "0.9.7" +version = "0.10.0" description = "LightAgent: Lightweight AI agent framework with memory, tools & tree-of-thought. Supports multi-agent collaboration, self-learning, and major LLMs (OpenAI/DeepSeek/Qwen). Open-source with MCP/SSE protocol integration." authors = ["caiweige "] license = "Apache-2.0" diff --git a/roadmap.md b/roadmap.md index 2e94c08..76ab1a0 100644 --- a/roadmap.md +++ b/roadmap.md @@ -1,6 +1,6 @@ # LightAgent Roadmap -Last updated: 2026-08-09 +Last updated: 2026-08-15 LightAgent should continue to evolve as a lightweight, low-dependency agent framework rather than a broad replacement for LangChain, LangGraph, CrewAI, or @@ -8,9 +8,9 @@ LlamaIndex. The product direction remains: -**Lightweight core + composable Skills + reliable tool execution + observable -traces + safe memory + deterministic workflows + OpenAI-compatible model -ecosystem.** +**Lightweight core + event-sourced Sessions + composable capability Providers + +reliable tool execution + unified Policy + safe memory + deterministic +workflows + OpenAI-compatible model ecosystem.** ## Current Status @@ -59,13 +59,17 @@ ecosystem.** evaluation, tool/handoff human review, durable LightFlow approvals, review batches, human feedback, and shared Graph Memory fail-closed write admission and audit controls. - -### In Development - - **v0.9.7**: Added the dependency-free Connector manifest and offline validation contract, two credential-free examples, expanded Python executor adversarial checks, an opt-in real Mem0 Graph security matrix, and the first - v1.0 public API compatibility inventory. Pending pull request and release. + public API compatibility inventory for the v1.0 stabilization line. + +### In Development + +- **v0.10.0**: Deliver the unified event-sourced Agent Runtime, combining + durable Sessions, native async execution, Capability Registry and Policy, + Inbox/Goals/Budgets, compaction and recovery, multi-Agent Jobs/Workflow, and + standardized Memory/Skills/MCP/RAG Providers while preserving v0.9.x APIs. ### Completed Milestone Details @@ -116,9 +120,9 @@ result = flow.run("Analyze this company") ### Open Pull Requests -- No open pull requests as of 2026-08-09. PR #85 was merged before v0.9.7 - development; v0.9.7 builds on it with broader adversarial and false-positive - regression coverage plus explicit execution-safety documentation. +- Live pull-request state changes faster than this roadmap and should be read + from GitHub. Documentation-only pull requests do not change the runtime + version plan or release gates recorded here. ### Active Issues @@ -137,10 +141,6 @@ P1 engineering work: retrieval-filter audit counts. The remaining acceptance criterion is an opt-in test against the exact Mem0 Graph version and storage configuration used in production. -- **#5 Custom plugin/integration development**: define a small connector - contract that can bundle Tools, Skills, MCP settings, Hooks, memory adapters, - optional dependencies, and docs without creating a heavy marketplace or - required plugin runtime. - **#1 Enhanced memory management for multi-agent systems**: keep shared-memory adapter hardening active until durable graph/vector backends have explicit tenant, provenance, conflict, and trust-boundary tests. @@ -154,6 +154,9 @@ P2 issues: Resolved or ready to close: +- **#5 Custom plugin/integration development**: v0.9.7 delivered the optional, + dependency-free Connector manifest, offline validator, examples, and + contributor documentation without adding a marketplace runtime. - **#33 Optional ClawMem memory backend**: #74 delivered the optional dependency-free adapter example, documentation, and fake-client tests. @@ -165,11 +168,11 @@ Not planned for the core repository: ## Near-Term Version Plan This section records the planned direction for the next several LightAgent -versions after `v0.9.6`. Exact scope can still change as issues, pull requests, +versions after `v0.9.7`. Exact scope can still change as issues, pull requests, and user feedback evolve, but the intended product direction is: -**security validation + lightweight connectors + safer execution tools + -stable APIs + production documentation.** +**event-sourced runtime + composable capability providers + unified policy + +long-task control + recoverable context + stable APIs.** ### v0.8.3 Goals: LightFlow Execution Controls @@ -585,7 +588,7 @@ humans in control of high-impact external side effects. ### v0.9.7: Security Validation, Connector Contract, And Release Hardening -Status: implemented and locally validated; pending pull request and release. +Status: released on 2026-08-15. Goal: close the remaining security and extensibility gaps before the v1.0 API freeze. v0.9.7 should be a bridge release: small enough to ship quickly, but @@ -647,69 +650,284 @@ Release gates: - Docs clearly distinguish built-in primitives from optional integration examples. -Local validation on the development branch: 193 passed, 1 opt-in Mem0 Graph -test skipped, package compilation and wheel build passed, and `git diff ---check` passed. Multi-version GitHub CI remains a pull-request release gate. +Release validation: 193 passed, 1 opt-in Mem0 Graph test skipped, package +compilation and wheel build passed, `git diff --check` passed, and GitHub CI +passed on Python 3.10, 3.11, 3.12, and 3.13. Expected outcome: -LightAgent should enter the v1.0 stabilization phase with fewer loose security +LightAgent enters the v0.10 runtime-evolution phase with fewer loose security threads, a practical answer to custom integrations, and clearer boundaries around what the lightweight core will and will not own. -### v1.0.0: Stable API And Production Documentation +### Post-v0.9.7 Runtime Design Guardrails + +The next development line should strengthen the runtime without turning the +core package into a hosted platform or a mandatory collection of heavyweight +integrations. + +- Preserve `agent.run("hello")`, `stream=True`, structured results, existing + Tools, Hooks, Memory backends, LightSwarm, and the LightFlow chain API. +- Make every model-visible message, tool result, memory item, approval result, + steering message, and compaction summary reconstructable from durable Session + events. +- Separate capability execution from policy decisions: a Provider implements + an operation, while Policy, Sandbox, Guardrails, and Approval decide whether + that operation may run. +- Require child Agents, Skills, and Workflow steps to inherit or reduce parent + permissions; task code must never expand its own capability set. +- Make the runtime async-first while retaining synchronous compatibility + wrappers. +- Keep Browser, Docker, LSP, vector databases, WebUI frameworks, and hosted + services as optional Providers or upper-layer product capabilities. +- Distinguish model errors, tool failures, policy denials, approval waits, + budget exhaustion, cancellation, and context overflow with explicit states + and error codes. + +### v0.10.0: Unified Event-Sourced Agent Runtime + +Status: implementation candidate completed on `codex/develop-v0.10.0`; +release validation is in progress. + +Goal: deliver one coherent runtime release that combines the previously +planned v0.10.0-v0.15.0 capabilities without breaking v0.9.x applications. +The work remains ordered as six internal milestones, but there are no separate +public v0.11.0-v0.15.0 releases in this plan. + +Implementation delivered in the v0.10.0 development PR: + +- Versioned Session events, in-memory/JSONL/SQLite stores, replay, pagination, + checkpoints, fork lineage, migration hooks, context/trace projection, and + explicit incomplete-Turn detection. +- Compatible `arun()`/`astream()` entry points and durable model/tool/runtime + lifecycle recording without changing default `run()` return behavior. +- Scoped Capability Registry, Provider lifecycle, deterministic conflict + diagnostics, narrowing-only permissions, unified Policy decisions, audit + configuration digests, and adapters for Tools, Memory, Skills, MCP, and + LightFlow. +- Durable Inbox, Goals, Budgets, progress detection, Jobs, bounded subagents, + context budgets, deterministic/summary compaction, oversized-tool spill + references, Session control events, and restart restoration. +- Deterministic Skill precedence, conflict reporting, nested `AGENTS.md` + discovery, MCP Streamable HTTP/reconnect/refresh/namespaces/credential + headers, SQLite FTS5 RAG, and citation-based cross-Session search. +- Focused v0.10 protocol, persistence, corruption, policy, runtime, async, + compatibility, and retrieval tests plus complete legacy regression testing. + +External validation still required before release: + +- Python 3.10-3.13 GitHub Actions, package build/install, and real provider + smoke tests. +- Fault injection against real MCP Streamable HTTP, process interruption, + concurrent persistent writes, and context-overflow provider responses. +- Contract tests for optional Browser, Terminal, Shell, LSP, vector, sandbox, + and hosted-service Providers supplied outside the lightweight core. + +#### Milestone 1: Event-Sourced Sessions And Native Async + +- Add versioned `Session`, `SessionEvent`, and `SessionStore` contracts. +- Define Session, Turn, Step, Message, Model, Tool, Approval, Error, and + lifecycle events with schema validation and migration hooks. +- Provide dependency-free in-memory and JSONL stores plus an optional SQLite + store based on the Python standard library. +- Add Session export, pagination, replay, recovery, and incomplete-Turn + detection. +- Derive model context and the current `TraceRecorder` view from the same event + history instead of maintaining unrelated sources of truth. +- Add native `agent.arun()` and retain `run()` as a compatibility wrapper. +- Record balanced model request/response and tool request/result pairs with + explicit interrupted and failed terminal states. + +#### Milestone 2: Capability Registry And Unified Policy + +- Add `CapabilityProvider` and `CapabilityRegistry` protocols with mount, + start, health, reload, stop, and unmount lifecycle methods. +- Support Runtime, Session, and Agent scopes with deterministic resolution and + conflict diagnostics. +- Define protocols for Model, Tool, FileSystem, Shell, Terminal, Browser, Web, + LSP, Memory, RAG, Subagent, Workflow, Interaction, Sandbox, Credential, + Policy, and Telemetry Providers. +- Adapt existing Tools, MCP, Memory, Connector, LightFlow, Hooks, Guardrails, + and approval APIs instead of introducing a parallel plugin runtime. +- Add capability metadata for read/write/network/execute behavior, risk, + timeout, output limits, cancellation, persistence, and optional dependencies. +- Route sensitive operations through one Policy decision path and record the + Provider name, version, and configuration digest in audit events. + +#### Milestone 3: Agent Inbox, Goals, And Budgets + +- Add a durable Agent Inbox for `followup`, `steering`, `context`, and + `approval` messages. +- Queue and consume messages in order, injecting steering only at safe Step + boundaries. +- Add durable Goals with acceptance criteria, subgoals, completion evidence, + blockers, and status transitions. +- Add model-call, tool-call, token, time, and estimated-cost budgets. +- Support pause, resume, cancel, and continue through Session events. +- Add no-progress detection, repeated-tool detection, bounded retry, and + message idempotency keys. + +#### Milestone 4: Context Compaction, Checkpoints, And Fork + +- Add model-aware token accounting and configurable context budgets. +- Implement two-stage compaction: deterministic trimming first, optional LLM + summarization second. +- Spill oversized tool results outside the prompt while retaining event-backed + references and integrity metadata. +- Persist compaction summaries and covered event ranges as versioned Session + events. +- Add Session checkpoints, restore validation, and Fork from a selected event + boundary. +- Support bounded recovery from context-overflow errors and an optional + dedicated summarization model. + +#### Milestone 5: Multi-Agent, Jobs, And Workflow Unification + +- Unify LightSwarm, handoff, and subagent lifecycle events while preserving + existing LightSwarm behavior. +- Support one-shot, persistent, and Session-Fork subagents with depth, count, + concurrency, and budget limits. +- Add Agent-tree inspection, messaging, interruption, resume, and result + collection. +- Freeze auditable child-permission snapshots and prohibit capability + escalation. +- Add background Jobs with status, incremental output, cancellation, and Inbox + completion notifications. +- Evolve LightFlow into the common Workflow Provider for fixed DAGs, dynamic + model-planned workflows, checkpoints, approvals, reruns, and parallel steps. +- Add optional persistent Terminal and LSP Providers without making them core + dependencies. -Goal: freeze the public API surface and make LightAgent dependable for -production users and contributors. +#### Milestone 6: Memory, Skills, MCP, And Knowledge Standardization + +- Standardize Working, Session, Workspace, User, and Shared Memory scopes with + owner, tenant, provenance, TTL, trust, sensitivity, and admission metadata. +- Keep automatic Memory writes and promotion behind `MemoryPolicy`, Policy, + and optional approval. +- Support user, workspace, nested-directory, managed, and built-in Markdown + Skills with deterministic precedence and conflict diagnostics. +- Add compatible project instruction discovery such as `AGENTS.md` without + runtime self-modification. +- Add MCP Streamable HTTP, reconnect, tool-list refresh, namespace isolation, + and external Credential Provider integration while retaining stdio/SSE + configuration compatibility. +- Define a Retrieval/RAG Provider and ship an optional SQLite FTS5 minimum + implementation; keep embeddings, vector databases, reranking, and hybrid + retrieval optional. +- Add cross-Session text search with citations while keeping Session Search + separate from knowledge-base retrieval. + +#### v0.10.0 Compatibility Commitments + +- Preserve `agent.run("hello")`, `stream=True`, structured results, existing + Tools, Hooks, Guardrails, Memory backends, LightSwarm, and LightFlow APIs. +- Existing users can adopt Session, Registry, Inbox, Goal, compaction, + subagent, and knowledge features incrementally; none becomes mandatory for a + basic Agent. +- Existing Trace, Tool, Memory, Hook, MCP, LightSwarm, and LightFlow data is + exposed through compatibility adapters instead of forced migration. +- Browser, Docker, LSP, vector databases, hosted services, and WebUI frameworks + remain optional. + +#### v0.10.0 Release Gates + +- Every model request can be reconstructed deterministically from persisted + Session events, and model/tool/approval records remain balanced. +- Process interruption, EventLog failure, context overflow, Provider failure, + and incomplete Turns have explicit recoverable or terminal states. +- Provider contract and cleanup tests prove that replacement and unload do not + leak tools, listeners, processes, credentials, or stale registrations. +- Write, network, execution, credential, and persistence operations cannot + bypass Policy, approval, scope inheritance, or audit handling. +- Restart preserves Inbox order, Goal state, pending approvals, budgets, + checkpoints, and idempotency markers. +- Compaction preserves unresolved Goals, approvals, decisions, file changes, + tool lineage, and replay integrity. +- Child Agent and Job failures cannot erase parent state; concurrent writes are + denied or serialized unless explicitly allowed. +- Workflow and Agentic Loop execution use the same Session, capability, + Policy, approval, budget, and recovery contracts. +- Memory, Skill, MCP, RAG, and Session Search retain source, owner, scope, and + provenance metadata; MCP reconnect cannot duplicate tools. +- The Runtime remains usable without vector, Browser, Docker, hosted service, + or model-gateway dependencies. +- The complete v0.9.7 compatibility suite passes on Python 3.10-3.13, together + with replay, migration, corruption, fault-injection, long-task, concurrency, + security, package-build, and import tests. + +### v1.0.0: Stable Runtime And Ecosystem + +Goal: freeze the runtime contracts only after they have survived multiple +pre-1.0 releases and fault-oriented validation. Planned work: -- Stabilize public APIs: - - `LightAgent` - - `LightSwarm` - - `LightFlow` - - `Skill` - - `ToolRegistry` - - `MemoryProtocol` - - `MemoryPolicy` - - `GuardrailDecision` - - `RunResult` - - `StreamEvent` -- Reduce breaking changes after the 1.0 line. -- Complete bilingual documentation for installation, tools, skills, memory, - MCP, Guardrails, Trace, LightSwarm, LightFlow, and production deployment. -- Add a complete example matrix covering basic agents, constructor tools, - runtime tools, memory, Skills, MCP, browser-use, OpenRouter, LiteLLM, local - LLMs, LightFlow, human approval, and error handling. -- Add stronger CI coverage for core runtime behavior. -- Automate PyPI release publishing and release notes. +- Freeze the public API, Provider protocols, Session event schemas, Policy + decisions, and compatibility adapters. +- Publish a versioned deprecation and migration policy with tooling for v0.9.x + Session, Trace, Tool, Memory, Hook, LightSwarm, and LightFlow users. +- Provide a Headless Runner, Python SDK, and optional JSON-RPC service surface. +- Publish official Provider templates and contract-test kits. +- Complete multilingual production documentation and the supported example + matrix. +- Add OpenTelemetry, Langfuse, and JSONL exporters through optional adapters. +- Establish performance, recovery, tool-call, multi-agent, and workflow + reliability benchmarks. +- Automate signed package build, PyPI publishing, release notes, and rollback + checks. -Expected outcome: +Release gates: -LightAgent 1.0 should provide a stable, documented, tested foundation for -building lightweight production agents. +- Public contracts have passed the complete v0.10.0 milestone suite and at + least one release-candidate or stabilization-patch compatibility cycle. +- Event schemas support forward migration and deterministic replay. +- Long-task interruption recovery passes in deterministic test environments. +- Multi-Agent, approval, compaction, MCP, and Provider lifecycle paths pass + fault-injection tests. +- Core installation does not require Browser, Docker, vector databases, model + gateway SDKs, or Web frameworks. +- The v0.9.x-to-v1.0 migration guide and compatibility suite are complete. ### v1.1.0: Enterprise Integration -Goal: make LightAgent easier to embed into internal systems and private -deployments. +Goal: make LightAgent easier to embed into internal systems after the runtime +contracts are stable. Planned work: -- Add stronger multi-tenant memory isolation examples and policy templates. -- Provide tool-level permission and audit patterns for production systems. -- Add deployment templates for Docker and service-style API wrappers. -- Improve model routing guidance for OpenAI-compatible endpoints, LiteLLM, - local inference servers, and private model gateways. -- Provide audit log export examples for trace, tool calls, guardrail blocks, - memory writes, and workflow steps. -- Add enterprise-oriented examples for customer service, data analysis, - internal knowledge assistants, and automated office workflows. - -Expected outcome: - -LightAgent should be easier to adopt inside enterprise systems without turning -the core framework into a large platform. +- Add multi-tenant policy templates and reference deployment profiles. +- Provide tool-level permission, credential, and audit patterns. +- Add optional Docker and service-wrapper deployment templates. +- Improve model routing guidance for compatible endpoints, LiteLLM, local + inference, and private gateways. +- Add enterprise examples without placing business workflows or hosted user + interfaces in the core package. + +### Unified v0.10.0 Quality Gates + +Every internal v0.10.0 milestone must extend, not replace, the following +validation layers. Passing an early milestone does not authorize releasing a +partial v0.10.0 as the final version: + +1. Protocol and state-machine unit tests. +2. Provider contract and resource-cleanup tests. +3. Session replay, migration, projection, and corruption tests. +4. Compatibility tests for all v0.9.7 public APIs. +5. Security tests for Policy, Sandbox, Approval, Credential, and scope + inheritance. +6. Fault injection for model streams, tools, stores, MCP, Providers, Jobs, and + subagents. +7. Long-task tests covering budgets, compaction, checkpoint, resume, and + idempotency. +8. Python 3.10, 3.11, 3.12, and 3.13 CI plus package build and import checks. + +Suggested release cadence: + +| Version | Theme | Suggested cycle | +| --- | --- | --- | +| v0.10.0 | Unified event-sourced Agent Runtime | 24-36 weeks, milestone-driven | +| v1.0.0 | API freeze and production hardening | After v0.10 stabilization gates | +| v1.1.0 | Optional enterprise integration | Post-v1.0 feedback-driven | ## Reference Directions From Other Agent Frameworks @@ -1032,7 +1250,7 @@ promotion decisions. ### v0.9.7 Workstream: Security Validation And Connector Contract -Status: implemented and locally validated; pending pull request and release. +Status: released on 2026-08-15. Goal: harden the remaining high-risk surfaces and define a minimal custom integration path before the v1.0 API freeze. @@ -1075,45 +1293,22 @@ the shared Graph Memory disclosure, and a small but useful extension path for domain integrations, while keeping v1.0 focused on stability instead of new surface area. -### v1.0.0 Workstream: Stable API And Ecosystem - -Goal: stabilize the public API and make LightAgent reliable for production users -and contributors. - -### Planned Work - -- Stabilize public APIs: - - `LightAgent` - - `LightSwarm` - - `LightFlow` - - `Skill` - - `ToolRegistry` - - `MemoryProtocol` - - `MemoryPolicy` - - `RunResult` -- Build a complete documentation site. -- Add a full example matrix: - - basic agent; - - constructor tools; - - runtime tools; - - memory; - - Skills; - - MCP; - - browser-use; - - OpenRouter; - - LiteLLM; - - local LLM; - - LightFlow; - - human approval. -- Add CI coverage for core runtime behavior. -- Automate PyPI release publishing. -- Add benchmarks for tool-call success rate, multi-turn completion, token cost, - latency, and recovery behavior. +### Post-v0.9.7 Runtime Workstreams -### Expected Outcome +The earlier plan split runtime evolution across v0.10.0-v0.15.0. These scopes +are now consolidated into one public **v0.10.0 Unified Event-Sourced Agent +Runtime** release with six ordered internal milestones: + +1. Session events, stores, replay, projection, and native async execution. +2. Capability Registry, Provider lifecycle, scopes, and unified Policy. +3. Durable Inbox, Goals, budgets, steering, and cancellation. +4. Context compaction, checkpoints, restore, and Session Fork. +5. Subagents, background Jobs, and Workflow/Agent Loop unification. +6. Standardized Memory, Skills, MCP, Retrieval, and RAG Providers. -LightAgent 1.0 should provide a stable, documented, tested foundation for -building lightweight production agents. +The detailed scope and release gates are maintained in the Near-Term Version +Plan. v1.0 is deferred until the complete v0.10.0 runtime has passed its +compatibility, replay, recovery, security, and stabilization gates. ## Longer-Term Directions @@ -1172,77 +1367,73 @@ building lightweight production agents. - Respond to the #39 advisory/CVE request, move reproduction and version-scoping details into a private security workflow, and avoid naming affected or fully patched versions until the shared Graph Memory test matrix is complete. +- Run the opt-in matrix against every maintained Mem0 Graph and storage + configuration before changing public remediation claims. ### Next P1 -- Merge and release the completed v0.9.7 security validation and connector - contract before v1.0. -- #39 shared Graph Memory backend-level validation, tenant/provenance tests, - and public/private advisory wording. -- Close #5 after the v0.9.7 Connector manifest, offline validator, and - dependency-free examples are merged. -- v1.0 public API inventory, compatibility contracts, deprecation policy, and - production documentation preparation. +- Start **v0.10.0 Unified Event-Sourced Agent Runtime** with the smallest + stable Session event model and compatibility adapters, then advance through + all six internal milestones under the same public version. +- Implement native `arun()` without changing `run()` or streaming behavior. +- Add in-memory, JSONL, and optional SQLite Session stores with replay and + incomplete-Turn recovery tests. +- Convert Trace into a projection of Session history while preserving current + Trace APIs and exporters. +- Continue #39 backend validation as an independent security release gate. ### P2 -- Durable memory-review queue examples that build on v0.9.5 promotion - candidates without becoming required core dependencies. -- External trace/audit adapters and production review-queue examples built on - the v0.9.6 exporter and approval contracts. -- Database-backed workflow and shared-memory adapters. -- Stronger idempotency and distributed execution controls for persistent - workflows. -- Focused external provider examples only when maintained outside the core - dependency set. +- Prepare the v0.10.0 Capability Registry milestone and Provider contract-test + fixtures in parallel, but do not route production execution through them + before Session invariants are stable. +- Add fault-injection fixtures for interrupted model streams, tool timeout, + EventLog write failure, and concurrent Session recovery. +- Keep external Provider examples focused, optional, credential-free in CI, + and outside the required core dependency set. ### P3 -- Visual trace UI. -- Distributed execution and durable worker coordination. +- Inbox, Goal, Budget, compaction, subagents, background Jobs, Workflow, MCP, + and RAG remain required v0.10.0 milestones and must land after their Session + and Provider prerequisites instead of accumulating in one unreviewable + change. +- Visual trace UI and distributed worker coordination remain upper-layer or + post-protocol work. ## Next Development Recommendation -After v0.9.7 is reviewed and released, the next development target should be -**v1.0.0 Stable API And Production Documentation**. The main runtime, -workflow, memory-safety, observability, evaluation, human-review, and Connector -primitives now exist; the next priority is freezing the documented public API -and completing production release automation. +The next development target is **v0.10.0 Unified Event-Sourced Agent Runtime**. +It includes the complete former v0.10.0-v0.15.0 scope. Implementation remains +milestone-ordered, but the public version is released only after all six +milestones and their combined quality gates pass. Reasoning: -- v0.9.0 covers checkpointed LightFlow runs, resume/rerun, approval nodes, - memory-safety controls, guardrail templates, and the shared-memory prototype. -- v0.9.3 completes stream tool-loop safety and consistent runtime hook closure. -- v0.9.4 completes tool schema diagnostics, `PolicyHook`, `on_handoff`, - LightSwarm runtime-context propagation, and full tracked-test CI. -- v0.9.5 adds explicit promotion candidates, promotion decisions, - `before_memory_promote` / `after_memory_promote`, promotion trace events, and - fail-closed tests for internal/shared memory safety. -- v0.9.6 adds trace summaries/exporters, deterministic evaluation, tool and - handoff review, durable LightFlow approvals, human feedback, and the first - fake-backend #39 cross-user graph-memory regression. -- Public GitHub state still shows #39 and #5 as open P1 issues. v0.9.7 - implements #5's narrowed Connector scope, while #39 remains open until the - exact maintained Mem0 Graph configurations run the opt-in matrix. -- Follow-up #39 work keeps durable shared-memory poisoning, provenance, and - multi-agent memory boundaries as active P1 concerns. -- The merged #85 hardening now has broader adversarial tests and explicit - documentation that AST filtering is not a complete sandbox. -- #5 is best addressed before v1.0 as a lightweight connector contract, not a - broad marketplace or new plugin runtime. -- Database-backed durability should stay optional so the core package remains - lightweight. - -Completed v0.9.7 implementation slice: - -1. Documented the #39 security response boundary and added a backend-specific - opt-in Graph Memory validation test. -2. Built on merged #85 and expanded Python executor adversarial blocklist and - safe false-positive coverage. -3. Added Python executor safety docs and recommended `PolicyHook` / Human Review - wrappers for high-risk deployments. -4. Defined a Connector manifest and offline validator without adding required - runtime dependencies. -5. Added two dependency-free Connector examples and a contributor guide. -6. Published the first v1.0 public API inventory and compatibility matrix. +- Trace, Hooks, review, Memory, LightFlow, and streaming currently record + related lifecycle data through different surfaces; one durable EventLog is + required before reliable resume and context reconstruction can be promised. +- Long-running Agent execution needs native async cancellation and recovery + semantics rather than additional wrappers around the current synchronous + loop. +- Capability Registry and Policy unification depend on stable Session identity, + event ordering, and audit records, so milestone 1 must precede milestone 2 + even though both ship in v0.10.0. +- The six-milestone v0.10.0 plan reduces the risk of freezing immature + contracts in v1.0 while keeping development increments independently + reviewable and testable. +- Optional stores and Providers preserve the lightweight core and let + LightWorker or other products supply Browser, Docker, WebUI, and business + workflow implementations. + +First v0.10.0 implementation slice: + +1. Publish versioned Session event dataclasses and an in-memory store. +2. Record one non-streaming Agent run as balanced Session, Turn, Model, Tool, + and terminal events. +3. Rebuild current Trace events and model context from that Session history. +4. Add JSONL persistence, replay, incomplete-Turn detection, and corruption + tests. +5. Add native `arun()` and prove `run()` plus `stream=True` compatibility. +6. Add optional SQLite storage only after the store contract passes the same + replay and migration suite. diff --git a/tests/test_lightflow.py b/tests/test_lightflow.py index ae8e928..05e1f84 100644 --- a/tests/test_lightflow.py +++ b/tests/test_lightflow.py @@ -1,3 +1,4 @@ +import asyncio import time from LightAgent import HookDecision, JsonLightFlowStore, LightFlow, LightFlowResult, RunResult @@ -34,6 +35,27 @@ def test_lightflow_runs_single_step_and_returns_object_result(): assert agent.calls[0]["kwargs"]["run_group_id"] == result.run_id +def test_lightflow_arun_preserves_result_and_does_not_block_event_loop(): + agent = FakeAgent("writer", ["done"]) + flow = LightFlow().step("write", agent=agent) + + async def scenario(): + marker = [] + + async def tick(): + await asyncio.sleep(0) + marker.append("tick") + + result, _ = await asyncio.gather(flow.arun("draft"), tick()) + return result, marker + + result, marker = asyncio.run(scenario()) + + assert result.success is True + assert result.content == "done" + assert marker == ["tick"] + + def test_lightflow_passes_dependency_outputs_to_later_steps(): research = FakeAgent("research", ["facts"]) writer = FakeAgent("writer", ["report"]) diff --git a/tests/test_mcp_client_manager.py b/tests/test_mcp_client_manager.py index 2106973..98197c0 100644 --- a/tests/test_mcp_client_manager.py +++ b/tests/test_mcp_client_manager.py @@ -1,4 +1,5 @@ import asyncio +from types import SimpleNamespace import pytest @@ -6,6 +7,33 @@ from LightAgent.tools import ToolRegistry +def mcp_tool(name="search_docs"): + return SimpleNamespace( + name=name, + description="Search documentation", + inputSchema={ + "type": "object", + "properties": {"query": {"type": "string", "title": "Query"}}, + "required": ["query"], + }, + ) + + +class FakeSession: + def __init__(self, tools=None, *, list_error=None, result="ok"): + self.tools = tools or [] + self.list_error = list_error + self.result = result + + async def list_tools(self): + if self.list_error: + raise self.list_error + return SimpleNamespace(tools=self.tools) + + async def call_tool(self, name, arguments): + return SimpleNamespace(content=[SimpleNamespace(text=self.result)]) + + def test_mcp_call_preserves_sanitized_failure_receipt(): manager = MCPClientManager( { @@ -41,3 +69,131 @@ async def fail_create_session(server_name, config): "message": "connection failed with Bearer [redacted] and [redacted]", } ] + + +def test_mcp_call_reconnects_and_succeeds_after_transient_failure(): + manager = MCPClientManager( + {"mcpServers": {"docs": {"command": "unused", "args": [], "reconnect_attempts": 1}}}, + ToolRegistry(), + ) + sessions = iter([ + FakeSession(list_error=ConnectionError("temporary disconnect")), + FakeSession([mcp_tool()], result="found"), + ]) + attempts = [] + + async def create_session(server_name, config): + attempts.append(server_name) + manager.session = next(sessions) + manager.server_sessions[server_name] = manager.session + + manager._create_session = create_session + + result = asyncio.run(manager.call_tool("search_docs", {"query": "runtime"}, target_server="docs")) + + assert result == {"server": "docs", "tool": "search_docs", "result": "found"} + assert attempts == ["docs", "docs"] + assert manager.last_mcp_errors[0]["error_type"] == "ConnectionError" + assert manager.server_sessions == {} + + +def test_mcp_refresh_replaces_namespaced_tools_without_duplicates(): + config = { + "mcpServers": { + "docs one": { + "command": "unused", + "args": [], + "namespace": "docs-one", + "namespace_tools": True, + }, + "docs-two": { + "command": "unused", + "args": [], + "namespace_tools": True, + }, + } + } + registry = ToolRegistry() + manager = MCPClientManager(config, registry) + sessions = { + "docs one": FakeSession([mcp_tool()]), + "docs-two": FakeSession([mcp_tool()]), + } + + async def create_session(server_name, server_config): + manager.session = sessions[server_name] + manager.server_sessions[server_name] = manager.session + + manager._create_session = create_session + + assert asyncio.run(manager.register_mcp_tool()) is True + assert set(registry.function_mappings) == {"docs-one__search_docs", "docs-two__search_docs"} + + sessions["docs one"] = FakeSession([mcp_tool(), mcp_tool("read_page")]) + assert asyncio.run(manager.refresh_tools("docs one")) is True + + names = [schema["function"]["name"] for schema in registry.openai_function_schemas] + assert set(names) == {"docs-one__search_docs", "docs-one__read_page", "docs-two__search_docs"} + assert len(names) == len(set(names)) + assert set(config["mcpServers"]) == {"docs one", "docs-two"} + + +def test_streamable_http_uses_async_credential_provider(monkeypatch): + from LightAgent import mcp_client_manager as mcp_module + + captured = {} + + class AsyncContext: + def __init__(self, value): + self.value = value + + async def __aenter__(self): + return self.value + + async def __aexit__(self, exc_type, exc, traceback): + return False + + class InitializedSession: + initialized = False + + async def initialize(self): + self.initialized = True + + def streamable_client(*, url, headers): + captured.update(url=url, headers=headers) + return AsyncContext(("reader", "writer", "session-id")) + + initialized = InitializedSession() + monkeypatch.setattr(mcp_module, "streamablehttp_client", streamable_client) + monkeypatch.setattr(mcp_module, "ClientSession", lambda *streams: AsyncContext(initialized)) + + async def credentials(server_name, config): + assert server_name == "remote" + return {"Authorization": "Bearer dynamic-token"} + + manager = MCPClientManager( + { + "mcpServers": { + "remote": { + "transport": "streamable-http", + "url": "https://mcp.example.test", + "headers": {"X-Client": "LightAgent"}, + } + } + }, + ToolRegistry(), + credential_provider=credentials, + ) + + async def scenario(): + await manager._create_session("remote", manager.config["mcpServers"]["remote"]) + assert manager.server_sessions["remote"] is initialized + await manager.cleanup() + + asyncio.run(scenario()) + + assert initialized.initialized + assert captured == { + "url": "https://mcp.example.test", + "headers": {"X-Client": "LightAgent", "Authorization": "Bearer dynamic-token"}, + } diff --git a/tests/test_skill_manager_logging.py b/tests/test_skill_manager_logging.py index 56aff0a..cf72723 100644 --- a/tests/test_skill_manager_logging.py +++ b/tests/test_skill_manager_logging.py @@ -1,5 +1,7 @@ import logging +import pytest + from LightAgent import LightAgent from LightAgent.logger import LoggerManager from LightAgent.skills import SkillManager @@ -78,3 +80,56 @@ def log(self, level, action, data): manager._log("DEBUG", "discover_skill", {"name": "demo"}) assert logger.calls == [("DEBUG", "discover_skill", {"name": "demo"})] + + +def test_later_skill_directory_wins_and_conflict_is_reported(tmp_path): + first = tmp_path / "first" + second = tmp_path / "second" + write_skill(first, description="First") + winner = write_skill(second, description="Second") + manager = SkillManager([str(first), str(second)]) + + manager.discover_skills() + + assert manager.skills["demo"].description == "Second" + assert manager.skills["demo"].path == str(winner) + assert manager.list_conflicts() == [{ + "name": "demo", + "winner": str(winner), + "shadowed": str(first / "demo"), + "rule": "later skills_directories entry wins", + }] + + +def test_project_instructions_are_loaded_root_to_leaf_and_bounded(tmp_path): + project = tmp_path / "project" + nested = project / "src" / "feature" + nested.mkdir(parents=True) + project.joinpath("AGENTS.md").write_text("project rules", encoding="utf-8") + nested.joinpath("AGENTS.md").write_text("feature rules", encoding="utf-8") + manager = SkillManager([]) + + instructions = manager.discover_project_instructions(nested) + + assert instructions.index("project rules") < instructions.index("feature rules") + with pytest.raises(ValueError, match="exceed max_chars"): + manager.discover_project_instructions(nested, max_chars=5) + + +def test_skill_reference_asset_and_script_paths_cannot_escape(tmp_path): + skills_dir = tmp_path / "skills" + skill_dir = write_skill(skills_dir) + skill_dir.joinpath("references").mkdir() + skill_dir.joinpath("references", "guide.md").write_text("safe", encoding="utf-8") + skill_dir.joinpath("assets").mkdir() + skill_dir.joinpath("assets", "data.bin").write_bytes(b"safe") + skill_dir.joinpath("scripts").mkdir() + manager = SkillManager([str(skills_dir)]) + manager.discover_skills() + + assert manager.read_reference("demo", "guide.md") == "safe" + assert manager.read_asset("demo", "data.bin") == b"safe" + assert "Security violation" in manager.read_reference("demo", "../../AGENTS.md") + with pytest.raises(ValueError, match="Security violation"): + manager.read_asset("demo", "../../outside.bin") + assert "Security violation" in manager.execute_script("demo", "../../outside.py") diff --git a/tests/test_v010_capabilities.py b/tests/test_v010_capabilities.py new file mode 100644 index 0000000..9f28788 --- /dev/null +++ b/tests/test_v010_capabilities.py @@ -0,0 +1,192 @@ +import asyncio + +import pytest + +from LightAgent import ( + BaseCapabilityProvider, + CapabilityRegistry, + CapabilityRisk, + CapabilityScope, + CapabilitySpec, + PermissionSet, + PolicyDecision, + PolicyEngine, + ProviderHealth, + RuntimeContext, +) + + +class EchoProvider(BaseCapabilityProvider): + name = "echo" + version = "test" + + def __init__(self, prefix=""): + super().__init__([CapabilitySpec("text.echo", read=True, timeout=0.1)]) + self.config = {"prefix": prefix, "api_key": "do-not-audit"} + self.prefix = prefix + + async def invoke(self, capability, **arguments): + return self.prefix + arguments["text"] + + +class ControlledProvider(BaseCapabilityProvider): + name = "controlled" + + def __init__(self, spec, result="result"): + super().__init__([spec]) + self.result = result + + async def invoke(self, capability, **arguments): + if arguments.get("wait"): + await asyncio.sleep(arguments["wait"]) + return self.result + + +def test_provider_lifecycle_and_health(): + provider = EchoProvider() + context = RuntimeContext(runtime_id="runtime") + + asyncio.run(provider.mount(context)) + asyncio.run(provider.start()) + health = asyncio.run(provider.health()) + asyncio.run(provider.stop()) + asyncio.run(provider.unmount()) + + assert health == ProviderHealth(healthy=True, status="ready") + assert not provider.mounted + + +def test_registry_scope_precedence_and_conflict_diagnostics(): + registry = CapabilityRegistry() + registry.register(EchoProvider("runtime:"), scope=CapabilityScope.RUNTIME) + registry.register(EchoProvider("agent:"), scope=CapabilityScope.AGENT, owner_id="a1") + + result = asyncio.run(registry.invoke( + "text.echo", {"text": "hello"}, context=RuntimeContext(agent_id="a1") + )) + + assert result == "agent:hello" + assert registry.conflicts() == [] + + +def test_equal_scope_conflict_is_diagnostic_and_latest_wins(): + registry = CapabilityRegistry() + registry.register(EchoProvider("one:")) + second = EchoProvider("two:") + second.name = "echo-two" + registry.register(second) + + assert registry.conflicts()[0]["winner"] == "echo-two" + assert asyncio.run(registry.invoke("text.echo", {"text": "x"})) == "two:x" + + +def test_permission_snapshot_cannot_escalate(): + parent = PermissionSet( + allowed=frozenset({"text.echo"}), + max_risk=CapabilityRisk.ISOLATED_WRITE, + ) + + child = parent.narrow(allowed={"text.echo"}, max_risk=CapabilityRisk.READ_ONLY) + + assert child.allows("text.echo") + with pytest.raises(ValueError, match="cannot add"): + parent.narrow(allowed={"text.echo", "shell.execute"}) + + +def test_policy_can_rewrite_or_deny_arguments(): + def policy(request): + if request.arguments["text"] == "deny": + return PolicyDecision.block("blocked by test") + return PolicyDecision.allow({"text": request.arguments["text"].upper()}) + + registry = CapabilityRegistry(policy_engine=PolicyEngine([policy])) + registry.register(EchoProvider()) + + assert asyncio.run(registry.invoke("text.echo", {"text": "ok"})) == "OK" + with pytest.raises(PermissionError, match="blocked by test"): + asyncio.run(registry.invoke("text.echo", {"text": "deny"})) + + +def test_audit_contains_digest_but_not_credentials(): + events = [] + registry = CapabilityRegistry(audit=lambda event, data: events.append((event, data))) + registry.register(EchoProvider()) + + asyncio.run(registry.invoke("text.echo", {"text": "ok"})) + + decision = next(data for event, data in events if event == "policy.decision") + assert len(decision["configuration_digest"]) == 64 + assert "do-not-audit" not in repr(decision) + + +@pytest.mark.parametrize("fail_closed,allowed", [(True, False), (False, True)]) +def test_policy_exception_respects_fail_closed_mode(fail_closed, allowed): + async def broken_policy(request): + await asyncio.sleep(0) + raise RuntimeError("policy backend unavailable") + + engine = PolicyEngine([broken_policy], fail_closed=fail_closed) + request = RuntimeContext() + registry = CapabilityRegistry(policy_engine=engine) + registry.register(EchoProvider()) + + if allowed: + assert asyncio.run(registry.invoke("text.echo", {"text": "ok"}, context=request)) == "ok" + else: + with pytest.raises(PermissionError, match="policy `broken_policy` failed: RuntimeError"): + asyncio.run(registry.invoke("text.echo", {"text": "ok"}, context=request)) + + +def test_capability_requires_approval_before_provider_invocation(): + provider = ControlledProvider(CapabilitySpec("dangerous.write", requires_approval=True)) + registry = CapabilityRegistry() + registry.register(provider) + + with pytest.raises(PermissionError, match="requires approval"): + asyncio.run(registry.invoke("dangerous.write")) + + +def test_capability_timeout_and_output_limit_are_enforced(): + timed = ControlledProvider(CapabilitySpec("slow.read", timeout=0.01)) + limited = ControlledProvider(CapabilitySpec("limited.read", output_limit=4), result="abcdefgh") + limited.name = "limited" + registry = CapabilityRegistry() + registry.register(timed) + registry.register(limited) + + with pytest.raises(asyncio.TimeoutError): + asyncio.run(registry.invoke("slow.read", {"wait": 0.05})) + assert asyncio.run(registry.invoke("limited.read")) == "abcd" + + +def test_registry_lifecycle_reload_unregister_and_reverse_stop_order(): + calls = [] + + class LifecycleProvider(EchoProvider): + async def stop(self): + calls.append(f"stop:{self.name}") + await super().stop() + + async def unmount(self): + calls.append(f"unmount:{self.name}") + await super().unmount() + + first = LifecycleProvider() + first.name = "first" + second = LifecycleProvider() + second.name = "second" + registry = CapabilityRegistry() + registry.register(first) + registry.register(second) + + async def scenario(): + await registry.mount(RuntimeContext(runtime_id="runtime")) + await registry.reload("first", {"mode": "strict"}) + assert await registry.unregister("missing") is False + assert await registry.unregister("second") is True + await registry.stop() + + asyncio.run(scenario()) + + assert first.config == {"mode": "strict"} + assert calls == ["stop:second", "unmount:second", "stop:first", "unmount:first"] diff --git a/tests/test_v010_core_integration.py b/tests/test_v010_core_integration.py new file mode 100644 index 0000000..a8896dc --- /dev/null +++ b/tests/test_v010_core_integration.py @@ -0,0 +1,230 @@ +import asyncio +import json +from types import SimpleNamespace + +from LightAgent import BudgetLimits, InMemorySessionStore, LightAgent + + +class StaticCompletions: + def __init__(self, replies): + self.replies = iter(replies) + self.requests = [] + + def create(self, **kwargs): + self.requests.append(kwargs) + message = SimpleNamespace(content=next(self.replies), tool_calls=None) + return SimpleNamespace( + choices=[SimpleNamespace(message=message)], + usage=SimpleNamespace(prompt_tokens=2, completion_tokens=1, total_tokens=3), + ) + + +class ToolCallCompletions: + def __init__(self): + self.requests = [] + + def create(self, **kwargs): + self.requests.append(kwargs) + if len(self.requests) == 1: + tool_call = SimpleNamespace( + id="call-add", + function=SimpleNamespace(name="add", arguments=json.dumps({"a": 2, "b": 3})), + ) + return SimpleNamespace( + choices=[SimpleNamespace(message=SimpleNamespace(content=None, tool_calls=[tool_call]))] + ) + return SimpleNamespace( + choices=[SimpleNamespace(message=SimpleNamespace(content="five", tool_calls=None))] + ) + + +def add(a, b): + return a + b + + +add.tool_info = { + "tool_name": "add", + "tool_description": "Add two numbers.", + "tool_params": [ + {"name": "a", "type": "number", "description": "First", "required": True}, + {"name": "b", "type": "number", "description": "Second", "required": True}, + ], +} + + +def make_agent(replies, **kwargs): + agent = LightAgent( + model="test-model", + api_key="test-key", + auto_discover_skills=False, + **kwargs, + ) + completions = StaticCompletions(replies) + agent.client = SimpleNamespace(chat=SimpleNamespace(completions=completions)) + return agent, completions + + +def test_run_persists_balanced_turn_and_exact_model_request(): + store = InMemorySessionStore() + agent, _ = make_agent(["hello"], session_store=store) + + assert agent.run("hi", session_id="session-1") == "hello" + + session = store.get("session-1") + event_types = [event.type for event in session.events] + assert event_types.count("turn.started") == 1 + assert event_types.count("turn.completed") == 1 + assert event_types.count("model.requested") == 1 + assert event_types.count("model.completed") == 1 + request = next(event for event in session.events if event.type == "model.requested") + assert request.data["messages"][-1] == {"role": "user", "content": "hi"} + assert session.replay().incomplete_turns == [] + + +def test_explicit_session_continues_previous_conversation(): + store = InMemorySessionStore() + agent, completions = make_agent(["first answer", "second answer"], session_store=store) + + agent.run("first", session_id="conversation") + agent.run("second", session_id="conversation") + + second_messages = completions.requests[1]["messages"] + assert {"role": "assistant", "content": "first answer"} in second_messages + assert second_messages[-1] == {"role": "user", "content": "second"} + + +def test_arun_preserves_legacy_string_result(): + agent, _ = make_agent(["async answer"]) + + result = asyncio.run(agent.arun("hello")) + + assert result == "async answer" + + +def test_astream_returns_async_iterator(): + agent, _ = make_agent(["unused"]) + + def fake_run(query, **kwargs): + return iter(["a", "b"]) + + agent.run = fake_run + + async def collect(): + stream = await agent.arun("hello", stream=True) + return [chunk async for chunk in stream] + + assert asyncio.run(collect()) == ["a", "b"] + + +def test_model_call_budget_blocks_before_second_request(): + agent, _ = make_agent(["first", "second"], budget_limits=BudgetLimits(model_calls=1)) + + assert agent.run("first", session_id="budgeted") == "first" + second = agent.run("second", session_id="budgeted") + + assert second.startswith("[LA-BUDGET]") + assert agent.replay_session()["failed_turns"] + + +def test_tool_call_persists_balanced_request_and_result_events(): + store = InMemorySessionStore() + agent = LightAgent( + model="test-model", + api_key="test-key", + auto_discover_skills=False, + session_store=store, + ) + completions = ToolCallCompletions() + agent.client = SimpleNamespace(chat=SimpleNamespace(completions=completions)) + + assert agent.run("add", tools=[add], session_id="tools") == "five" + + session = store.get("tools") + requested = [event for event in session.events if event.type == "tool.requested"] + completed = [event for event in session.events if event.type == "tool.completed"] + assert len(requested) == len(completed) == 1 + assert requested[0].data["name"] == completed[0].data["name"] == "add" + assert requested[0].turn_id == completed[0].turn_id + assert requested[0].run_id == completed[0].run_id + assert session.replay().incomplete_turns == [] + + +def test_model_exception_persists_failed_model_run_and_turn(): + class FailingCompletions: + def create(self, **kwargs): + raise RuntimeError("provider unavailable") + + store = InMemorySessionStore() + agent, _ = make_agent(["unused"], session_store=store) + agent.client = SimpleNamespace(chat=SimpleNamespace(completions=FailingCompletions())) + + result = agent.run("hello", session_id="failed", result_format="object") + + session = store.get("failed") + types = [event.type for event in session.events] + assert result.error + assert types.count("model.requested") == 1 + assert types.count("model.failed") == 1 + assert types.count("run.failed") == 1 + assert types.count("turn.failed") == 1 + assert session.replay().failed_turns + assert session.replay().incomplete_turns == [] + + +def test_astream_closes_sync_iterator_when_consumer_stops_early(): + closed = [] + + def stream_values(): + try: + yield "first" + yield "second" + finally: + closed.append(True) + + agent, _ = make_agent(["unused"]) + agent.run = lambda query, **kwargs: stream_values() + + async def consume_one(): + stream = await agent.arun("hello", stream=True) + assert await anext(stream) == "first" + await stream.aclose() + + asyncio.run(consume_one()) + + assert closed == [True] + + +def test_public_session_control_apis_persist_events_and_fork(): + store = InMemorySessionStore() + agent, _ = make_agent(["answer"], session_store=store) + agent.run("hello", session_id="control") + + checkpoint = agent.checkpoint_session("before review") + agent.pause_session("review") + agent.resume_session("approved") + agent.cancel_session("finished") + forked = agent.fork_session(through_sequence=checkpoint["sequence"]) + + session = store.get("control") + assert [event.type for event in session.events[-4:]] == [ + "session.checkpointed", "session.paused", "session.resumed", "session.cancelled" + ] + assert forked.metadata["forked_from"] == "control" + assert store.get(forked.session_id) is not None + + +def test_public_compact_session_reduces_projected_history_and_persists_event(): + store = InMemorySessionStore() + agent, _ = make_agent(["one", "two", "three"], session_store=store) + agent.run("first", session_id="compact") + agent.run("second", session_id="compact") + agent.run("third", session_id="compact") + + result = agent.compact_session(max_messages=2) + + assert result["removed_count"] > 0 + assert len(result["messages"]) <= 2 + persisted = store.get("compact") + event = persisted.events[-1] + assert event.type == "context.compacted" + assert event.data["removed_count"] == result["removed_count"] diff --git a/tests/test_v010_knowledge.py b/tests/test_v010_knowledge.py new file mode 100644 index 0000000..726caa8 --- /dev/null +++ b/tests/test_v010_knowledge.py @@ -0,0 +1,193 @@ +import asyncio + +import pytest + +from LightAgent import ( + CapabilityRegistry, + InMemorySessionStore, + MCPProviderAdapter, + RetrievalDocument, + RuntimeContext, + Session, + SessionSearchProvider, + SkillProviderAdapter, + SqliteFTSRetrievalProvider, + WorkflowProviderAdapter, +) +from LightAgent.skills import Skill + + +def test_sqlite_fts_ingest_search_citation_and_scope(tmp_path): + provider = SqliteFTSRetrievalProvider(tmp_path / "rag.sqlite3", chunk_size=100, chunk_overlap=10) + provider.ingest(RetrievalDocument( + document_id="d1", + title="Runtime", + source="guide.md", + content="LightAgent durable session checkpoint recovery", + scope="workspace", + owner_id="team-a", + )) + + results = provider.search("checkpoint", owner_id="team-a") + + assert results[0].document_id == "d1" + assert results[0].citation_id.startswith("rag:d1:") + assert provider.search("checkpoint", owner_id="team-b") == [] + + +def test_rag_provider_can_be_invoked_through_registry(tmp_path): + provider = SqliteFTSRetrievalProvider(tmp_path / "rag.sqlite3") + registry = CapabilityRegistry() + registry.register(provider) + + document_id = asyncio.run(registry.invoke("rag.ingest", { + "document": {"content": "provider registry policy", "document_id": "d1"} + }, context=RuntimeContext(runtime_id="r1"))) + results = asyncio.run(registry.invoke("rag.search", {"query": "registry"})) + + assert document_id == "d1" + assert results[0].document_id == "d1" + + +def test_session_search_is_separate_and_returns_event_citation(): + store = InMemorySessionStore() + session = Session(session_id="s1") + session.append("message.received", {"content": "quarterly forecast"}) + store.create(session) + provider = SessionSearchProvider(store) + + result = provider.search("forecast")[0] + + assert result["citation_id"] == "session:s1:1" + assert result["event_type"] == "message.received" + + +def test_rag_document_lifecycle_chunking_and_reindex(tmp_path): + provider = SqliteFTSRetrievalProvider(tmp_path / "rag.sqlite3", chunk_size=100, chunk_overlap=20) + document = RetrievalDocument( + document_id="long", + title="Long document", + source="long.md", + content="checkpoint " + ("x" * 130) + " recovery", + metadata={"version": 1}, + ) + + assert provider.ingest(document) == "long" + chunks = [provider.read_chunk(f"long:{index}") for index in range(2)] + assert [len(chunk.content) for chunk in chunks] == [100, 70] + assert chunks[0].metadata == {"version": 1} + assert provider.read_chunk("missing") is None + assert [item.document_id for item in provider.list_documents()] == ["long"] + + provider.reindex() + assert provider.search("checkpoint")[0].document_id == "long" + assert provider.remove("long") is True + assert provider.remove("long") is False + assert provider.list_documents() == [] + assert provider.read_chunk("long:0") is None + + +def test_rag_scope_owner_and_tenant_filters_are_combined(tmp_path): + provider = SqliteFTSRetrievalProvider(tmp_path / "rag.sqlite3") + provider.ingest(RetrievalDocument( + document_id="tenant-a", + content="shared runtime handbook", + scope="project", + owner_id="owner-a", + tenant_id="tenant-a", + )) + provider.ingest(RetrievalDocument( + document_id="tenant-b", + content="shared runtime handbook", + scope="project", + owner_id="owner-b", + tenant_id="tenant-b", + )) + + matches = provider.search( + "runtime", scope="project", owner_id="owner-a", tenant_id="tenant-a", limit=10 + ) + + assert [result.document_id for result in matches] == ["tenant-a"] + assert provider.search("runtime", owner_id="owner-a", tenant_id="tenant-b") == [] + assert provider.search(" ") == [] + with pytest.raises(ValueError, match="limit must be at least 1"): + provider.search("runtime", limit=0) + + +def test_rag_like_fallback_reports_degraded_health_and_searches(tmp_path): + provider = SqliteFTSRetrievalProvider(tmp_path / "rag.sqlite3") + provider.ingest(RetrievalDocument(document_id="fallback", content="Fallback Search Works")) + provider._fts5 = False + + health = asyncio.run(provider.health()) + + assert health.status == "degraded" + assert health.degraded_capabilities == ["rag.search"] + assert provider.search("search")[0].document_id == "fallback" + provider.reindex() # A degraded provider keeps the source tables intact. + + +def test_session_search_is_case_insensitive_bounded_and_handles_empty_query(): + store = InMemorySessionStore() + for index in range(3): + session = Session(session_id=f"s{index}") + session.append("message.received", {"content": f"Quarterly Forecast {index}"}) + store.create(session) + provider = SessionSearchProvider(store) + + matches = provider.search("FORECAST", limit=2) + + assert len(matches) == 2 + assert all(match["citation_id"].startswith("session:") for match in matches) + assert provider.search(" ") == [] + + +def test_skill_mcp_and_workflow_provider_adapters_delegate_arguments(): + class Skills: + def __init__(self): + self.skills = {"demo": Skill(name="demo", description="Demo", path="skills/demo")} + + def activate_skill(self, skill_name): + return f"active:{skill_name}" + + class MCP: + async def call_tool(self, name, arguments): + return {"name": name, "arguments": arguments} + + class Flow: + async def arun(self, query): + return f"started:{query}" + + def get_run(self, run_id): + return {"run_id": run_id} + + def resume(self, run_id, **kwargs): + return {"resumed": run_id, **kwargs} + + def rerun_step(self, run_id, step_name): + return {"rerun": run_id, "step": step_name} + + skills = SkillProviderAdapter(Skills()) + mcp = MCPProviderAdapter(MCP(), ["search"]) + workflow = WorkflowProviderAdapter(Flow()) + + async def scenario(): + listed = await skills.invoke("skill.list") + activated = await skills.invoke("skill.activate", skill_name="demo") + called = await mcp.invoke("mcp.search", query="runtime") + started = await workflow.invoke("workflow.start", query="draft") + status = await workflow.invoke("workflow.status", run_id="run-1") + resumed = await workflow.invoke("workflow.resume", run_id="run-1", trace=True) + rerun = await workflow.invoke("workflow.rerun_step", run_id="run-1", step_name="write") + return listed, activated, called, started, status, resumed, rerun + + listed, activated, called, started, status, resumed, rerun = asyncio.run(scenario()) + + assert listed[0]["name"] == "demo" + assert activated == "active:demo" + assert called == {"name": "search", "arguments": {"query": "runtime"}} + assert started == "started:draft" + assert status == {"run_id": "run-1"} + assert resumed == {"resumed": "run-1", "trace": True} + assert rerun == {"rerun": "run-1", "step": "write"} diff --git a/tests/test_v010_runtime.py b/tests/test_v010_runtime.py new file mode 100644 index 0000000..832e079 --- /dev/null +++ b/tests/test_v010_runtime.py @@ -0,0 +1,215 @@ +import asyncio + +import pytest + +from LightAgent import ( + AgentRuntime, + BudgetExceeded, + BudgetLimits, + GoalStatus, + InboxMessageStatus, + InMemorySessionStore, + JobStatus, + PermissionSet, + ProgressTracker, +) + + +def test_inbox_goal_and_budget_restore_from_session(): + store = InMemorySessionStore() + runtime = AgentRuntime(session_store=store, budget_limits=BudgetLimits(model_calls=3)) + runtime.open_session("durable") + message = runtime.inbox.enqueue("steering", "focus", message_id="stable-id") + runtime.inbox.claim_next(safe_boundary=True) + runtime.inbox.complete(message.message_id) + goal = runtime.goals.create("ship", acceptance_criteria=["tests pass"]) + runtime.goals.activate(goal.goal_id) + runtime.budget.consume(model_calls=2) + + restored = AgentRuntime(session_store=store, budget_limits=BudgetLimits(model_calls=3)) + restored.open_session("durable") + + assert restored.inbox.get("stable-id").status == InboxMessageStatus.COMPLETED + assert restored.goals.get(goal.goal_id).status == GoalStatus.ACTIVE + assert restored.budget.remaining()["model_calls"] == 1 + assert restored.inbox.enqueue("steering", "duplicate", message_id="stable-id").content == "focus" + + +def test_budget_is_fail_closed(): + runtime = AgentRuntime(budget_limits=BudgetLimits(tool_calls=1)) + runtime.open_session() + runtime.budget.consume(tool_calls=1) + + with pytest.raises(BudgetExceeded): + runtime.budget.consume(tool_calls=1) + assert any(event.type == "budget.exhausted" for event in runtime.session.events) + + +def test_progress_tracker_detects_repeated_tools(): + tracker = ProgressTracker(max_repeated_tool_calls=2) + for _ in range(3): + tracker.record(tool="search", arguments={"q": "same"}) + + assert tracker.stalled + + +def test_background_job_reports_output_and_completion_to_inbox(): + async def scenario(): + runtime = AgentRuntime() + runtime.open_session() + + async def work(): + await asyncio.sleep(0) + return "done" + + record = runtime.jobs.start("work", work) + runtime.jobs.emit_output(record.job_id, "halfway") + completed = await runtime.jobs.wait(record.job_id) + return runtime, completed + + runtime, completed = asyncio.run(scenario()) + + assert completed.status == JobStatus.SUCCESS + assert completed.output == ["halfway"] + assert runtime.inbox.pending()[0].correlation_id == completed.job_id + + +class StubAgent: + name = "child" + + async def arun(self, query, **kwargs): + return query.upper() + + +def test_subagent_permissions_are_frozen_and_tree_is_inspectable(): + runtime = AgentRuntime() + runtime.open_session() + parent = PermissionSet(allowed=frozenset({"text.echo"})) + record = runtime.subagents.register( + StubAgent(), parent_permissions=parent, allowed_capabilities={"text.echo"} + ) + + result = asyncio.run(runtime.subagents.run(record.agent_id, "hello")) + + assert result == "HELLO" + assert runtime.subagents.tree()[0]["status"] == "success" + + +def test_steering_waits_for_safe_boundary_and_rejection_is_restored(): + store = InMemorySessionStore() + runtime = AgentRuntime(session_store=store) + runtime.open_session("inbox") + message = runtime.inbox.enqueue("steering", "change direction") + + assert runtime.inbox.claim_next(safe_boundary=False) is None + claimed = runtime.inbox.claim_next(safe_boundary=True) + assert claimed.message_id == message.message_id + runtime.inbox.reject(message.message_id, "unsafe request") + + restored = AgentRuntime(session_store=store) + restored.open_session("inbox") + assert restored.inbox.get(message.message_id).status == InboxMessageStatus.REJECTED + + +def test_goal_terminal_state_and_evidence_are_restored(): + store = InMemorySessionStore() + runtime = AgentRuntime(session_store=store) + runtime.open_session("goals") + goal = runtime.goals.create("release", acceptance_criteria=["tests pass"]) + runtime.goals.activate(goal.goal_id) + runtime.goals.complete(goal.goal_id, evidence=[{"suite": "passed"}]) + + with pytest.raises(ValueError, match="terminal goal"): + runtime.goals.block(goal.goal_id, "too late") + + restored = AgentRuntime(session_store=store) + restored.open_session("goals") + completed = restored.goals.get(goal.goal_id) + assert completed.status == GoalStatus.COMPLETED + assert completed.evidence == [{"suite": "passed"}] + + +@pytest.mark.parametrize("dimension,value", [ + ("model_calls", 2), + ("tool_calls", 2), + ("tokens", 11), + ("seconds", 1.5), + ("cost", 0.6), +]) +def test_each_budget_dimension_fails_without_committing_usage(dimension, value): + runtime = AgentRuntime(budget_limits=BudgetLimits(**{dimension: value / 2})) + runtime.open_session() + + with pytest.raises(BudgetExceeded) as exc_info: + runtime.budget.consume(**{dimension: value}) + + assert exc_info.value.dimension == dimension + assert getattr(runtime.budget.usage, dimension) == 0 + assert runtime.session.events[-1].type == "budget.exhausted" + + +def test_background_job_failure_is_persisted_and_reported_to_inbox(): + async def scenario(): + runtime = AgentRuntime() + runtime.open_session() + + async def fail(): + raise RuntimeError("job failed") + + record = runtime.jobs.start("failure", fail) + return runtime, await runtime.jobs.wait(record.job_id) + + runtime, failed = asyncio.run(scenario()) + + assert failed.status == JobStatus.FAILED + assert failed.error == "RuntimeError: job failed" + assert runtime.inbox.pending()[0].metadata["kind"] == "job_completion" + assert any(event.type == "job.failed" for event in runtime.session.events) + + +def test_background_job_cancellation_and_interrupted_restore(): + store = InMemorySessionStore() + + async def scenario(): + runtime = AgentRuntime(session_store=store) + runtime.open_session("jobs") + started = asyncio.Event() + + async def wait_forever(): + started.set() + await asyncio.Event().wait() + + record = runtime.jobs.start("cancelled", wait_forever) + await started.wait() + assert runtime.jobs.cancel(record.job_id) + return runtime, await runtime.jobs.wait(record.job_id) + + runtime, cancelled = asyncio.run(scenario()) + assert cancelled.status == JobStatus.CANCELLED + + pending_event = runtime.session.append("job.created", { + "job": {**cancelled.to_dict(), "job_id": "interrupted", "status": "running"} + }) + assert pending_event.type == "job.created" + store.save(runtime.session) + restored = AgentRuntime(session_store=store) + restored.open_session("jobs") + assert restored.jobs.get("interrupted").status == JobStatus.INTERRUPTED + + +def test_runtime_pause_resume_cancel_checkpoint_and_fork_are_durable(): + store = InMemorySessionStore() + runtime = AgentRuntime(session_store=store) + runtime.open_session("control") + runtime.pause("manual review") + checkpoint = runtime.checkpoint("reviewed") + runtime.resume("approved") + runtime.cancel("operator stop") + forked = runtime.fork(through_sequence=checkpoint.sequence) + + persisted_types = [event.type for event in store.get("control").events] + assert persisted_types[-4:] == [ + "session.paused", "session.checkpointed", "session.resumed", "session.cancelled" + ] + assert forked.metadata["forked_from"] == "control" + assert forked.metadata["forked_at_sequence"] == checkpoint.sequence diff --git a/tests/test_v010_session.py b/tests/test_v010_session.py new file mode 100644 index 0000000..595ad76 --- /dev/null +++ b/tests/test_v010_session.py @@ -0,0 +1,197 @@ +import json +import sqlite3 + +import pytest + +from LightAgent import ( + ContextBudget, + ContextCompactor, + ContextProjector, + InMemorySessionStore, + JsonlSessionStore, + Session, + SessionEvent, + SessionMigrationRegistry, + SqliteSessionStore, +) + + +@pytest.mark.parametrize("store_factory", [ + lambda path: InMemorySessionStore(), + lambda path: JsonlSessionStore(path / "jsonl"), + lambda path: SqliteSessionStore(path / "sessions.sqlite3"), +]) +def test_session_store_round_trip_and_replay(tmp_path, store_factory): + store = store_factory(tmp_path) + session = Session(session_id="session-1", metadata={"api_key": "secret"}) + session.append("turn.started", {}, turn_id="turn-1") + session.append("message.received", {"role": "user", "content": "hello"}, turn_id="turn-1") + session.append("turn.completed", {}, turn_id="turn-1") + store.create(session) + + restored = store.get("session-1") + + assert restored is not None + assert restored.metadata["api_key"] == "[redacted]" + assert restored.replay().completed_turns == ["turn-1"] + assert [event.sequence for event in restored.events] == [1, 2, 3] + + +def test_incomplete_turn_and_corrupt_sequence_are_detected(): + session = Session(session_id="session-1") + session.append("turn.started", {}, turn_id="turn-1") + assert session.replay().incomplete_turns == ["turn-1"] + + payload = session.to_dict() + payload["events"][0]["sequence"] = 2 + with pytest.raises(ValueError, match="sequence gap"): + Session.from_dict(payload) + + +def test_fork_preserves_projectable_events_without_reusing_ids(): + session = Session(session_id="source") + source = session.append("message.received", {"role": "user", "content": "hello"}, turn_id="t1") + forked = session.fork(new_session_id="fork") + + assert ContextProjector().messages(forked) == [{"role": "user", "content": "hello"}] + imported = next(event for event in forked.events if event.type == "message.received") + assert imported.event_id != source.event_id + assert imported.data["source_event_id"] == source.event_id + + +def test_exact_model_request_can_be_reconstructed(): + session = Session(session_id="session-1") + messages = [{"role": "system", "content": "policy"}, {"role": "user", "content": "hello"}] + event = session.append("model.requested", {"messages": messages}) + + assert ContextProjector().model_request(session, sequence=event.sequence) == messages + + +def test_compaction_spills_large_tool_results_and_meets_budget(): + messages = [ + {"role": "system", "content": "policy"}, + {"role": "user", "content": "first"}, + {"role": "assistant", "content": "working"}, + {"role": "tool", "content": "x" * 500}, + {"role": "user", "content": "last"}, + ] + compactor = ContextCompactor(max_inline_tool_chars=100) + result = compactor.compact_to_budget( + messages, + ContextBudget(max_tokens=100, reserved_output_tokens=10, chars_per_token=4), + ) + + assert result.spilled[0]["sha256"] + assert result.spilled[0]["content"] == "x" * 500 + assert all(len(message.get("content", "")) < 500 for message in result.messages) + + +def test_migration_registry_applies_ordered_event_migration(): + registry = SessionMigrationRegistry() + registry.register_event(0, lambda value: {**value, "schema_version": 1, "type": "message.received"}) + payload = SessionEvent(type="legacy", session_id="s").to_dict() + payload["schema_version"] = 0 + + migrated = registry.migrate_event(payload) + + assert migrated["schema_version"] == 1 + assert migrated["type"] == "message.received" + + +def test_jsonl_store_rejects_truncated_event(tmp_path): + store = JsonlSessionStore(tmp_path) + session = store.create(Session(session_id="broken")) + path = tmp_path / "broken.jsonl" + path.write_text(path.read_text(encoding="utf-8") + "{broken\n", encoding="utf-8") + + with pytest.raises(json.JSONDecodeError): + store.get(session.session_id) + + +@pytest.mark.parametrize("store_factory", [ + lambda path: InMemorySessionStore(), + lambda path: JsonlSessionStore(path / "jsonl"), + lambda path: SqliteSessionStore(path / "sessions.sqlite3"), +]) +def test_session_store_list_delete_and_duplicate_create(tmp_path, store_factory): + store = store_factory(tmp_path) + store.create(Session(session_id="one")) + store.create(Session(session_id="two")) + + assert {session.session_id for session in store.list(limit=2)} == {"one", "two"} + with pytest.raises(ValueError, match="already exists"): + store.create(Session(session_id="one")) + assert store.delete("one") is True + assert store.delete("one") is False + assert store.get("one") is None + + +def test_session_page_append_event_and_future_schema_validation(): + session = Session(session_id="session") + first = session.append("one") + session.append_event(SessionEvent(type="two", session_id="session")) + + assert [event.type for event in session.page(after=first.sequence, limit=1)] == ["two"] + with pytest.raises(ValueError, match="limit must be at least 1"): + session.page(limit=0) + with pytest.raises(ValueError, match="does not match"): + session.append_event(SessionEvent(type="bad", session_id="other")) + + payload = session.to_dict() + payload["events"][0]["schema_version"] = 999 + with pytest.raises(ValueError, match="unsupported SessionEvent schema_version"): + Session.from_dict(payload) + + +def test_in_memory_store_rejects_stale_session_save(): + store = InMemorySessionStore() + original = Session(session_id="session") + original.append("first") + stale = store.create(original) + current = store.get("session") + current.append("second") + store.save(current) + + with pytest.raises(ValueError, match="older event sequence"): + store.save(stale) + + +def test_jsonl_atomic_save_failure_preserves_previous_session(tmp_path, monkeypatch): + from LightAgent import session as session_module + + store = JsonlSessionStore(tmp_path) + current = store.create(Session(session_id="atomic")) + current.append("durable.event", {"value": 1}) + store.save(current) + before = store.get("atomic").to_dict() + current.append("new.event", {"value": 2}) + + def fail_replace(source, destination): + raise OSError("simulated replace failure") + + monkeypatch.setattr(session_module.os, "replace", fail_replace) + with pytest.raises(OSError, match="simulated replace failure"): + store.save(current) + + assert store.get("atomic").to_dict() == before + + +def test_sqlite_store_rejects_corrupt_persisted_event(tmp_path): + path = tmp_path / "sessions.sqlite3" + store = SqliteSessionStore(path) + store.create(Session(session_id="corrupt")) + with sqlite3.connect(path) as connection: + connection.execute( + "INSERT INTO session_events(session_id, sequence, event_id, payload) VALUES (?, ?, ?, ?)", + ("corrupt", 1, "bad-event", "{not-json"), + ) + + with pytest.raises(json.JSONDecodeError): + store.get("corrupt") + + +def test_jsonl_store_rejects_unsafe_session_id(tmp_path): + store = JsonlSessionStore(tmp_path) + + with pytest.raises(ValueError, match="unsafe characters"): + store.create(Session(session_id="../escape"))