1"""Callback manager implementation for AI observability."""
2
3from __future__ import annotations
4
5from typing import Any
6
7from lexigram.contracts.ai.callbacks import (
8 CallbackHandlerProtocol,
9)
10from lexigram.contracts.ai.llm import ChatMessage, Completion
11from lexigram.logging import (
12 get_logger,
13)
14
15logger = get_logger(__name__)
16
17
18class CallbackManagerImpl:
19 """Implementation of CallbackManagerProtocol.
20
21 Manages callback handlers and fans out events to all registered
22 handlers. Supports parent-child hierarchy for run-scoped propagation.
23 Handler failures are isolated to prevent breaking other handlers.
24 """
25
26 def __init__(self, parent: CallbackManagerImpl | None = None) -> None:
27 """Initialize the callback manager.
28
29 Args:
30 parent: Optional parent manager for child creation.
31 """
32 self._parent = parent
33 self._handlers: list[CallbackHandlerProtocol] = []
34 if parent is not None:
35 self._handlers = list(parent._handlers)
36
37 def register(self, handler: CallbackHandlerProtocol) -> None:
38 """Register a callback handler.
39
40 Args:
41 handler: The handler to register.
42 """
43 if handler not in self._handlers:
44 self._handlers.append(handler)
45
46 def unregister(self, handler: CallbackHandlerProtocol) -> None:
47 """Unregister a callback handler.
48
49 Args:
50 handler: The handler to remove.
51 """
52 if handler in self._handlers:
53 self._handlers.remove(handler)
54
55 def child(self, run_id: str) -> CallbackManagerImpl:
56 """Create a child callback manager for a specific run.
57
58 Args:
59 run_id: Unique identifier for this run.
60
61 Returns:
62 A child CallbackManagerImpl instance.
63 """
64 return CallbackManagerImpl(parent=self)
65
66 async def on_llm_start(
67 self,
68 messages: list[ChatMessage],
69 model: str,
70 **kwargs: Any,
71 ) -> None:
72 """Called when an LLM call starts."""
73 await self._fanout("on_llm_start", messages=messages, model=model, **kwargs)
74
75 async def on_llm_new_token(
76 self,
77 token: str,
78 **kwargs: Any,
79 ) -> None:
80 """Called for each new token in a streaming LLM response."""
81 await self._fanout("on_llm_new_token", token=token, **kwargs)
82
83 async def on_llm_end(
84 self,
85 response: Completion,
86 **kwargs: Any,
87 ) -> None:
88 """Called when an LLM call completes successfully."""
89 await self._fanout("on_llm_end", response=response, **kwargs)
90
91 async def on_llm_error(
92 self,
93 error: Exception,
94 **kwargs: Any,
95 ) -> None:
96 """Called when an LLM call fails."""
97 await self._fanout("on_llm_error", error=error, **kwargs)
98
99 async def on_chain_start(
100 self,
101 name: str,
102 inputs: dict[str, Any],
103 **kwargs: Any,
104 ) -> None:
105 """Called when a chain/pipeline starts executing."""
106 await self._fanout("on_chain_start", name=name, inputs=inputs, **kwargs)
107
108 async def on_chain_end(
109 self,
110 name: str,
111 outputs: dict[str, Any],
112 **kwargs: Any,
113 ) -> None:
114 """Called when a chain/pipeline completes."""
115 await self._fanout("on_chain_end", name=name, outputs=outputs, **kwargs)
116
117 async def on_tool_start(
118 self,
119 tool_name: str,
120 arguments: dict[str, Any],
121 **kwargs: Any,
122 ) -> None:
123 """Called when a tool starts executing."""
124 await self._fanout(
125 "on_tool_start", tool_name=tool_name, arguments=arguments, **kwargs
126 )
127
128 async def on_tool_end(
129 self,
130 tool_name: str,
131 result: Any,
132 **kwargs: Any,
133 ) -> None:
134 """Called when a tool finishes executing."""
135 await self._fanout("on_tool_end", tool_name=tool_name, result=result, **kwargs)
136
137 async def on_agent_action(
138 self,
139 action: dict[str, Any],
140 **kwargs: Any,
141 ) -> None:
142 """Called when an agent takes an action."""
143 await self._fanout("on_agent_action", action=action, **kwargs)
144
145 async def on_agent_finish(
146 self,
147 response: dict[str, Any],
148 **kwargs: Any,
149 ) -> None:
150 """Called when an agent finishes executing."""
151 await self._fanout("on_agent_finish", response=response, **kwargs)
152
153 async def on_retriever_start(
154 self,
155 query: str,
156 **kwargs: Any,
157 ) -> None:
158 """Called when a retriever starts a search."""
159 await self._fanout("on_retriever_start", query=query, **kwargs)
160
161 async def on_retriever_end(
162 self,
163 documents: list[Any],
164 **kwargs: Any,
165 ) -> None:
166 """Called when a retriever completes a search."""
167 await self._fanout("on_retriever_end", documents=documents, **kwargs)
168
169 async def _fanout(
170 self,
171 method_name: str,
172 **kwargs: Any,
173 ) -> None:
174 """Fan out a callback to all registered handlers.
175
176 Failures in individual handlers are isolated to prevent
177 breaking the entire chain.
178
179 Args:
180 method_name: Name of the callback method to invoke.
181 **kwargs: Arguments to pass to the callback.
182 """
183 for handler in self._handlers:
184 try:
185 method = getattr(handler, method_name)
186 await method(**kwargs)
187 except Exception as e:
188 logger.warning(
189 "Callback handler failed",
190 handler=handler.__class__.__name__,
191 method=method_name,
192 error=str(e),
193 )
194
195
196__all__ = ["CallbackManagerImpl"]