77 lines
2.1 KiB
Python
77 lines
2.1 KiB
Python
"""Runtime context for tool construction."""
|
|
from __future__ import annotations
|
|
|
|
from contextlib import contextmanager
|
|
from contextvars import ContextVar, Token
|
|
from dataclasses import dataclass, field
|
|
from typing import TYPE_CHECKING, Any, Callable, Protocol, runtime_checkable
|
|
|
|
if TYPE_CHECKING:
|
|
from nanobot.utils.llm_runtime import LLMRuntime
|
|
|
|
_CURRENT_REQUEST_CONTEXT: ContextVar["RequestContext | None"] = ContextVar(
|
|
"nanobot_tool_request_context",
|
|
default=None,
|
|
)
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class RequestContext:
|
|
"""Per-request context injected into tools at message-processing time."""
|
|
channel: str
|
|
chat_id: str
|
|
message_id: str | None = None
|
|
session_key: str | None = None
|
|
original_user_text: str | None = None
|
|
runtime: LLMRuntime | None = None
|
|
metadata: dict[str, Any] = field(default_factory=dict)
|
|
|
|
|
|
@runtime_checkable
|
|
class ContextAware(Protocol):
|
|
def set_context(self, ctx: RequestContext) -> None:
|
|
...
|
|
|
|
|
|
def bind_request_context(ctx: RequestContext) -> Token[RequestContext | None]:
|
|
return _CURRENT_REQUEST_CONTEXT.set(ctx)
|
|
|
|
|
|
def reset_request_context(token: Token[RequestContext | None]) -> None:
|
|
_CURRENT_REQUEST_CONTEXT.reset(token)
|
|
|
|
|
|
@contextmanager
|
|
def request_context(ctx: RequestContext):
|
|
"""Bind one immutable request snapshot and restore the previous value."""
|
|
token = bind_request_context(ctx)
|
|
try:
|
|
yield ctx
|
|
finally:
|
|
reset_request_context(token)
|
|
|
|
|
|
def current_request_context() -> RequestContext | None:
|
|
return _CURRENT_REQUEST_CONTEXT.get()
|
|
|
|
|
|
def current_request_session_key() -> str | None:
|
|
ctx = current_request_context()
|
|
return ctx.session_key if ctx else None
|
|
|
|
|
|
@dataclass
|
|
class ToolContext:
|
|
config: Any
|
|
workspace: str
|
|
bus: Any | None = None
|
|
subagent_manager: Any | None = None
|
|
cron_service: Any | None = None
|
|
sessions: Any | None = None
|
|
file_state_store: Any = field(default=None)
|
|
provider_snapshot_loader: Callable[[], Any] | None = None
|
|
image_generation_provider_configs: dict[str, Any] | None = None
|
|
timezone: str = "UTC"
|
|
workspace_sandbox: Any | None = None
|
|
runtime_events: Any | None = None
|