by-framework-langgraph 0.0.3.dev0__py3-none-any.whl → 0.0.3.dev2__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -16,7 +16,8 @@ from by_framework.common.logger import logger
16
16
  from by_framework.core.protocol.agent_state import AgentState
17
17
  from by_framework.core.protocol.commands import ResumeCommand
18
18
  from by_framework.core.protocol.events import StreamChunkEvent
19
- from by_framework.observability.span_recorder import (str_to_uint64, str_to_uint128)
19
+ from by_framework.trace.span_recorder import str_to_uint64, str_to_uint128
20
+ from langchain_core.callbacks import BaseCallbackHandler
20
21
  from langchain_core.messages import HumanMessage
21
22
  from langgraph.types import Command
22
23
 
@@ -31,6 +32,172 @@ if TYPE_CHECKING:
31
32
  LANGFUSE_OBSERVATION_ATTR = "_langfuse_observation"
32
33
 
33
34
 
35
+ class _TokenAccumulatingCallbackHandler(BaseCallbackHandler):
36
+ """LangChain callback handler that accumulates LLM token usage into AgentContext.
37
+
38
+ Extracts token usage from whichever location the provider populates:
39
+ - ``llm_output["token_usage"]`` (OpenAI-style via LangChain)
40
+ - ``llm_output["usage"]`` (Anthropic-style / raw provider mapping)
41
+ - ``generation.message.usage_metadata`` (LangChain >= 0.2 standard)
42
+ - ``generation.generation_info`` (some community integrations)
43
+
44
+ ``run_id`` deduplication prevents double-counting when both ``on_llm_end``
45
+ and ``on_chat_model_end`` fire for the same call (LangChain >= 0.2).
46
+ """
47
+
48
+ def __init__(self, context: Any) -> None:
49
+ super().__init__()
50
+ self._context = context
51
+ self._seen_run_ids: set = set()
52
+
53
+ # ------------------------------------------------------------------
54
+ # LangChain callback entry point
55
+ # ------------------------------------------------------------------
56
+
57
+ def on_llm_end(self, response: Any, *, run_id: Any = None, **_kwargs: Any) -> None:
58
+ # Guard: only mark run as seen when we actually extracted tokens so that
59
+ # the on_chat_model_end event path in _stream_invoke can still fire as a
60
+ # fallback when the callback found nothing (e.g. stream_options not set).
61
+ if run_id is not None and run_id in self._seen_run_ids:
62
+ return
63
+ self._handle_llm_result(response, run_id=run_id)
64
+
65
+ # ------------------------------------------------------------------
66
+ # Internal extraction logic
67
+ # ------------------------------------------------------------------
68
+
69
+ def _handle_llm_result(self, response: Any, *, run_id: Any = None) -> None:
70
+ context = self._context
71
+ if context is None:
72
+ return
73
+
74
+ prompt, completion = self._extract_tokens(response)
75
+
76
+ if not (prompt or completion):
77
+ # Log at WARNING (always visible) to help diagnose providers whose
78
+ # token format is not yet handled, or where stream_options is missing.
79
+ gens = getattr(response, "generations", []) or []
80
+ first = (gens[0] or [None])[0] if gens else None
81
+ msg_type = type(getattr(first, "message", None)).__name__
82
+ logger.warning(
83
+ "[TokenAccumulator] on_llm_end fired but extracted 0 tokens. "
84
+ "For OpenAI-compatible streaming APIs add "
85
+ "stream_options={'include_usage': True} to your ChatModel. "
86
+ "llm_output=%r first_gen_message_type=%s",
87
+ getattr(response, "llm_output", None),
88
+ msg_type,
89
+ )
90
+ # Do NOT mark run_id as seen — let the on_chat_model_end event path
91
+ # in _stream_invoke attempt extraction from the merged message.
92
+ return
93
+
94
+ if run_id is not None:
95
+ self._seen_run_ids.add(run_id)
96
+ try:
97
+ context.record_token_usage(
98
+ prompt_tokens=prompt,
99
+ completion_tokens=completion,
100
+ )
101
+ except Exception: # pylint: disable=broad-exception-caught
102
+ pass
103
+
104
+ @staticmethod
105
+ def _extract_tokens(response: Any) -> tuple[int, int]:
106
+ """Return (prompt_tokens, completion_tokens) from an LLMResult.
107
+
108
+ Checks every known location across providers:
109
+ 1. llm_output["token_usage"] — OpenAI via LangChain
110
+ 2. llm_output["usage"] — Anthropic / raw provider mapping
111
+ 3. message.usage_metadata — LangChain >= 0.2 standard
112
+ 4. message.response_metadata — some community integrations
113
+ 5. generation.generation_info — older / custom integrations
114
+ """
115
+ prompt, completion = 0, 0
116
+
117
+ llm_output = getattr(response, "llm_output", None) or {}
118
+ if isinstance(llm_output, dict):
119
+ for key in ("token_usage", "usage"):
120
+ usage = llm_output.get(key) or {}
121
+ if usage:
122
+ prompt = int(
123
+ usage.get("prompt_tokens") or usage.get("input_tokens") or 0
124
+ )
125
+ completion = int(
126
+ usage.get("completion_tokens")
127
+ or usage.get("output_tokens")
128
+ or 0
129
+ )
130
+ break
131
+
132
+ if prompt or completion:
133
+ return prompt, completion
134
+
135
+ # Iterate all generations
136
+ for gen_list in getattr(response, "generations", []) or []:
137
+ for gen in (gen_list if isinstance(gen_list, list) else [gen_list]):
138
+ msg = getattr(gen, "message", None)
139
+
140
+ # LangChain >= 0.2: message.usage_metadata
141
+ meta = getattr(msg, "usage_metadata", None)
142
+ if meta:
143
+ prompt += int(
144
+ meta.get("input_tokens") or meta.get("prompt_tokens") or 0
145
+ )
146
+ completion += int(
147
+ meta.get("output_tokens") or meta.get("completion_tokens") or 0
148
+ )
149
+ continue
150
+
151
+ # response_metadata (e.g. MiniMax, Qwen, some Chinese providers)
152
+ resp_meta = getattr(msg, "response_metadata", None) or {}
153
+ if isinstance(resp_meta, dict):
154
+ for key in ("token_usage", "usage"):
155
+ usage = resp_meta.get(key) or {}
156
+ if usage:
157
+ prompt += int(
158
+ usage.get("prompt_tokens")
159
+ or usage.get("input_tokens")
160
+ or 0
161
+ )
162
+ completion += int(
163
+ usage.get("completion_tokens")
164
+ or usage.get("output_tokens")
165
+ or 0
166
+ )
167
+ break
168
+ # Flat keys at root of response_metadata
169
+ if not (prompt or completion):
170
+ prompt += int(
171
+ resp_meta.get("prompt_tokens")
172
+ or resp_meta.get("input_tokens")
173
+ or 0
174
+ )
175
+ completion += int(
176
+ resp_meta.get("completion_tokens")
177
+ or resp_meta.get("output_tokens")
178
+ or 0
179
+ )
180
+ if prompt or completion:
181
+ continue
182
+
183
+ # generation_info fallback
184
+ info = getattr(gen, "generation_info", None) or {}
185
+ for key in ("token_usage", "usage"):
186
+ sub = info.get(key) or {}
187
+ if sub:
188
+ prompt += int(
189
+ sub.get("prompt_tokens") or sub.get("input_tokens") or 0
190
+ )
191
+ completion += int(
192
+ sub.get("completion_tokens")
193
+ or sub.get("output_tokens")
194
+ or 0
195
+ )
196
+ break
197
+
198
+ return prompt, completion
199
+
200
+
34
201
  @dataclass(frozen=True)
35
202
  class _AdapterTracingConfig:
36
203
  """Tracing-related adapter config kept separate from core graph handles."""
@@ -167,6 +334,47 @@ class LangGraphAdapter:
167
334
  await self._context.emit_chunk(
168
335
  chunk.content, content_type="text"
169
336
  )
337
+ elif kind == "on_chat_model_end":
338
+ # Fallback: capture token usage from the event's merged output
339
+ # message when the on_llm_end callback found nothing (e.g. the
340
+ # provider requires stream_options but it was not set).
341
+ # _TokenAccumulatingCallbackHandler marks run_id as seen only
342
+ # after a successful extraction, so this path fires only when
343
+ # the callback got 0 tokens.
344
+ run_id = event.get("run_id")
345
+ token_handler = next(
346
+ (
347
+ cb
348
+ for cb in scoped_callbacks
349
+ if isinstance(cb, _TokenAccumulatingCallbackHandler)
350
+ ),
351
+ None,
352
+ )
353
+ if token_handler is not None and run_id not in (
354
+ token_handler._seen_run_ids # pylint: disable=protected-access
355
+ ):
356
+ output = event["data"].get("output")
357
+ meta = getattr(output, "usage_metadata", None)
358
+ if meta:
359
+ prompt = int(
360
+ meta.get("input_tokens")
361
+ or meta.get("prompt_tokens")
362
+ or 0
363
+ )
364
+ completion = int(
365
+ meta.get("output_tokens")
366
+ or meta.get("completion_tokens")
367
+ or 0
368
+ )
369
+ if prompt or completion:
370
+ token_handler._seen_run_ids.add(run_id) # pylint: disable=protected-access
371
+ try:
372
+ self._context.record_token_usage(
373
+ prompt_tokens=prompt,
374
+ completion_tokens=completion,
375
+ )
376
+ except Exception: # pylint: disable=broad-exception-caught
377
+ pass
170
378
  elif kind == "on_tool_start":
171
379
  tool_name = event["name"]
172
380
  tool_input = event["data"].get("input")
@@ -184,7 +392,7 @@ class LangGraphAdapter:
184
392
  ),
185
393
  )
186
394
 
187
- # Use a stable logical ID from metadata if available, fallback to run_id
395
+ # Use stable metadata when available, otherwise fall back.
188
396
  stable_id = (
189
397
  event.get("metadata", {}).get("tool_call_id")
190
398
  or event.get("metadata", {}).get("checkpoint_ns")
@@ -379,69 +587,50 @@ class LangGraphAdapter:
379
587
  @contextmanager
380
588
  def _langfuse_callback_manager(self, callbacks: list[Any]) -> Iterator[None]:
381
589
  """Prepare Langfuse callback and observation for LangChain."""
382
- # Prefer AgentContext's callback factory so trace and parent ids align.
383
- # Filter out auto-generated MagicMock attributes when tests use a mock
384
- # context — real callback objects always come from a non-test module.
385
- langfuse_callback_value = getattr(self._context, "langfuse_callback", None)
386
- is_real_callback = langfuse_callback_value is not None and type(
387
- langfuse_callback_value
388
- ).__module__ not in ("unittest.mock",)
389
-
390
- if is_real_callback:
391
- handler = (
392
- langfuse_callback_value()
393
- if callable(langfuse_callback_value)
394
- else langfuse_callback_value
395
- )
396
- if handler is not None:
397
- callbacks.append(handler)
398
- yield
399
- return
590
+ # Always inject token accumulator — works regardless of Langfuse config.
591
+ callbacks.append(_TokenAccumulatingCallbackHandler(self._context))
400
592
 
401
- # Fallback to local import if context method is missing
402
- # pylint: disable=import-outside-toplevel
403
593
  try:
404
- langfuse_config = import_module(
405
- "by_framework_trace_langfuse"
406
- ).LangfuseConfig
407
- if langfuse_config.from_env() is None:
408
- raise ImportError("Langfuse not configured")
409
-
410
- callback_handler = import_module("langfuse.langchain").CallbackHandler
411
- get_client = import_module("langfuse").get_client
594
+ build_langchain_callback = getattr(
595
+ import_module("by_framework_trace_langfuse"),
596
+ "build_langchain_callback",
597
+ )
412
598
  except (ImportError, AttributeError):
413
599
  yield
414
600
  return
415
601
 
416
- callbacks.append(callback_handler())
417
-
418
- framework_observation = getattr(self._context, LANGFUSE_OBSERVATION_ATTR, None)
419
- if framework_observation is None:
420
- yield
421
- return
602
+ get_parent_observation_id = getattr(
603
+ self._context,
604
+ "get_trace_parent_observation_id",
605
+ None,
606
+ )
607
+ parent_observation_id = (
608
+ str(get_parent_observation_id() or "")
609
+ if callable(get_parent_observation_id)
610
+ else ""
611
+ )
612
+ if not parent_observation_id:
613
+ framework_observation = getattr(
614
+ self._context, LANGFUSE_OBSERVATION_ATTR, None
615
+ )
616
+ parent_observation_id = getattr(framework_observation, "id", "") or ""
617
+ if not parent_observation_id:
618
+ execution_id = getattr(self._context, "execution_id", "")
619
+ message_id = getattr(self._context, "message_id", "")
620
+ raw_parent_id = (
621
+ f"{execution_id}:worker.execute"
622
+ if execution_id
623
+ else f"{message_id}:worker.execute"
624
+ )
625
+ parent_observation_id = f"{str_to_uint64(raw_parent_id):016x}"
422
626
 
423
- langfuse = get_client()
424
- with langfuse.start_as_current_observation(
425
- as_type="span",
426
- name=self._tracing.run_name,
427
- trace_context={
428
- "trace_id": getattr(self._context, "trace_id", ""),
429
- "parent_span_id": framework_observation.id,
430
- },
431
- metadata=self._default_metadata(),
432
- ):
433
- # Prevent the generated OTel span from being promoted to a trace root.
434
- # The native LangfusePlugin sets the same attribute on its own path
435
- # (via _SdkLangfuseTracer); this covers the LangGraph fallback path.
436
- try:
437
- from opentelemetry import trace
438
-
439
- current_span = trace.get_current_span()
440
- if current_span and hasattr(current_span, "set_attribute"):
441
- current_span.set_attribute("langfuse.internal.as_root", False)
442
- except Exception: # pylint: disable=broad-exception-caught
443
- pass
444
- yield
627
+ handler = build_langchain_callback(
628
+ trace_id=getattr(self._context, "trace_id", ""),
629
+ parent_observation_id=parent_observation_id,
630
+ )
631
+ if handler is not None:
632
+ callbacks.append(handler)
633
+ yield
445
634
 
446
635
  @staticmethod
447
636
  def _default_input_mapper(content: str) -> dict:
@@ -7,7 +7,7 @@ mechanism.
7
7
 
8
8
  from __future__ import annotations
9
9
 
10
- from typing import TYPE_CHECKING, Annotated
10
+ from typing import TYPE_CHECKING, Annotated, Any
11
11
 
12
12
  from langchain_core.tools import BaseTool, InjectedToolCallId, tool
13
13
  from langgraph.types import interrupt
@@ -16,6 +16,31 @@ if TYPE_CHECKING:
16
16
  from by_framework.worker.context import AgentContext
17
17
 
18
18
 
19
+ def _langfuse_observation_id_from_callbacks(callbacks: Any) -> str:
20
+ """Return the Langfuse observation id for the active LangChain tool run."""
21
+ run_id = getattr(callbacks, "run_id", None) or getattr(
22
+ callbacks,
23
+ "parent_run_id",
24
+ None,
25
+ )
26
+ if not run_id:
27
+ return ""
28
+
29
+ handlers = [
30
+ *list(getattr(callbacks, "handlers", []) or []),
31
+ *list(getattr(callbacks, "inheritable_handlers", []) or []),
32
+ ]
33
+ for handler in handlers:
34
+ runs = getattr(handler, "_runs", None)
35
+ if not isinstance(runs, dict):
36
+ continue
37
+ observation = runs.get(run_id)
38
+ observation_id = getattr(observation, "id", None)
39
+ if observation_id:
40
+ return str(observation_id)
41
+ return ""
42
+
43
+
19
44
  def make_remote_agent_tool(
20
45
  context: AgentContext,
21
46
  tool_name: str,
@@ -50,6 +75,7 @@ def make_remote_agent_tool(
50
75
  async def remote_agent_tool(
51
76
  topic: str,
52
77
  tool_call_id: Annotated[str, InjectedToolCallId],
78
+ callbacks: Any = None,
53
79
  ) -> str:
54
80
  # Idempotency guard: checkpoint restore replays tool execution,
55
81
  # but we must not re-dispatch the command.
@@ -57,9 +83,19 @@ def make_remote_agent_tool(
57
83
  is_dispatched = await context.redis.exists(redis_key)
58
84
 
59
85
  if not is_dispatched:
86
+ metadata = {}
87
+ langfuse_parent_observation_id = _langfuse_observation_id_from_callbacks(
88
+ callbacks
89
+ )
90
+ if langfuse_parent_observation_id:
91
+ metadata["langfuse_parent_observation_id"] = (
92
+ langfuse_parent_observation_id
93
+ )
94
+
60
95
  await context.call_agent(
61
96
  target_agent_type=target_agent_type,
62
97
  content=topic,
98
+ metadata=metadata,
63
99
  )
64
100
  await context.redis.set(redis_key, "1", ex=idempotency_ttl)
65
101
 
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: by-framework-langgraph
3
- Version: 0.0.3.dev0
3
+ Version: 0.0.3.dev2
4
4
  Summary: LangGraph integration for by-framework
5
5
  Requires-Python: >=3.12
6
6
  Requires-Dist: by-framework>=0.2.1
@@ -0,0 +1,8 @@
1
+ by_framework_langgraph/__init__.py,sha256=x08wIaK4VRnUnUD_mbBxVgI2Zf7CvwAVbr3syP2NQps,1046
2
+ by_framework_langgraph/_utils.py,sha256=A8eQ19YqyBiMKcPIlzmwggN_mK5iQvNJPTaTmsywWsw,1605
3
+ by_framework_langgraph/adapter.py,sha256=b4WcK6R3B0Y7a52lIY6kMLtpp9mjXOkp4p6HcALAzbA,26689
4
+ by_framework_langgraph/tools.py,sha256=l7LurtKKXA1Z4fyj15ffFWyWpt8drNCJ-mkGmlS7EVo,5561
5
+ by_framework_langgraph/worker.py,sha256=wFalpBQLqFm-YzJugzAzkA3ciI9lUzxRTv4GmqVCv78,6156
6
+ by_framework_langgraph-0.0.3.dev2.dist-info/METADATA,sha256=lzuVDa2l8CUp_af7-VnadrCq-wJgXCWJ2ldcuEwHKRw,2334
7
+ by_framework_langgraph-0.0.3.dev2.dist-info/WHEEL,sha256=mffPy8wBnZQn2VnJUU5jE99KsxaSfiyMHV9Yt0aLVxs,87
8
+ by_framework_langgraph-0.0.3.dev2.dist-info/RECORD,,
@@ -1,8 +0,0 @@
1
- by_framework_langgraph/__init__.py,sha256=x08wIaK4VRnUnUD_mbBxVgI2Zf7CvwAVbr3syP2NQps,1046
2
- by_framework_langgraph/_utils.py,sha256=A8eQ19YqyBiMKcPIlzmwggN_mK5iQvNJPTaTmsywWsw,1605
3
- by_framework_langgraph/adapter.py,sha256=TmCU0gniDZxAXwnm-hA0M6wLaNThZ3CX3UznfG1g9-w,18051
4
- by_framework_langgraph/tools.py,sha256=T7OoNKMPuJT3t6bR5G8KVGSjrSnLP0tiCEzCnBGesMY,4384
5
- by_framework_langgraph/worker.py,sha256=wFalpBQLqFm-YzJugzAzkA3ciI9lUzxRTv4GmqVCv78,6156
6
- by_framework_langgraph-0.0.3.dev0.dist-info/METADATA,sha256=UI2zpgc9eB1ZBqy5BdNNqtonWWc-8QkmkpiE2dpE9MU,2334
7
- by_framework_langgraph-0.0.3.dev0.dist-info/WHEEL,sha256=mffPy8wBnZQn2VnJUU5jE99KsxaSfiyMHV9Yt0aLVxs,87
8
- by_framework_langgraph-0.0.3.dev0.dist-info/RECORD,,