mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-07-29 12:09:31 +08:00
feat(agent): report accurate aggregated token usage and propagate session/user + input/output to Langfuse for agent runs (#16420)
### What problem does this PR solve?
_Briefly describe what this PR aims to solve. Include background context
that will help reviewers understand the purpose of the PR._
### Type of change
- [x] Bug Fix (non-breaking change which fixes an issue)
- [x] New Feature (non-breaking change which adds functionality)
- [x] Other (please describe):
## Summary
Agent (Canvas) runs previously did not surface token usage in the SSE
stream, and RAGFlow's own Langfuse generations for agent runs were
missing the prompt/completion split and the session/user correlation.
This made it impossible for an external caller (or Langfuse) to
reconcile an agent turn's cost with the upstream provider (e.g.
OpenRouter), because a single turn can issue several distinct LLM calls
(query rewriting / cross-language translation, multi-round tool
reasoning, nested sub-agents, and the final answer).
This PR introduces a per-run token usage sink so that **every** LLM call
in a run is aggregated and reported once, and enriches Langfuse
generations with the prompt/completion split plus session/user
attributes.
## What changes
### 1. Per-run token usage sink (`common/token_utils.py`)
- Adds two `contextvars`: `token_usage_sink` (a mutable per-run
accumulator) and `langfuse_run_attrs` (session_id/user_id for the run).
- Adds `record_run_token_usage(...)` (thread-safe via a lock, because
`thread_pool_exec` copies the context into worker threads that share the
sink dict) and `usage_from_response(...)` which extracts a
`{prompt_tokens, completion_tokens, total_tokens}` split from
OpenAI/OpenRouter-style responses.
### 2. Provider layer captures the prompt/completion split
(`rag/llm/chat_model.py`)
- `LiteLLMBase` and `Base` now store `self.last_usage`
(prompt/completion/total) for the most recent chat call, in both the
plain and tool-calling paths.
- Streaming requests set `stream_options.include_usage = True` (LiteLLM
path) so the authoritative usage arrives on the final chunk; this is
read even on the usage-only chunk that carries no `choices`.
- Fixes a multi-round accounting bug in `*_with_tools`: token totals
were **overwritten** by each round (`total_tokens = tol`) instead of
accumulated, undercounting multi-round tool conversations. Each round is
now committed to a running aggregate.
### 3. LLMBundle reports usage once, per call
(`api/db/services/llm_service.py`)
- New `_report_usage(total_tokens)` records the call's usage into the
active run sink and returns the prompt/completion/total split for
Langfuse. The split is only used when it is consistent with the
authoritative total; otherwise only the total is reported.
- All three chat entry points (`async_chat`, `async_chat_streamly`,
`async_chat_streamly_delta`) now emit `usage_details` with
`input`/`output`/`total` instead of total-only.
- `_start_langfuse_observation` now applies `session_id`/`user_id` from
the per-run context (`langfuse_run_attrs`) so agent-run generations are
correctly grouped, even though agent LLMBundles are constructed without
those attributes.
### 4. Canvas installs the sink and emits the aggregate
(`agent/canvas.py`)
- `Canvas.run()` installs a fresh `token_usage_sink` and
`langfuse_run_attrs` (from `user_id`/`session_id`) at the start of every
turn.
- `message_end` now includes an aggregated `usage` object:
`{prompt_tokens, completion_tokens, total_tokens, calls}` covering all
LLM calls in the run.
### 5. Pass session id into the run
(`api/db/services/canvas_service.py`)
- `completion()` forwards `session_id` to `Canvas.run()` for Langfuse
session correlation.
## Why a context variable
LLM calls in an agent run originate from many places that each build
their own `LLMBundle` (e.g. `cross_languages`/`keyword_extraction`
helpers, the Agent component, and nested sub-agents invoked as tools). A
run-scoped context variable is the only non-invasive chokepoint that
captures all of them exactly once, including nested agents (which run in
the same async context) and thread-pool tools (the executor copies the
context).
## Behavior / compatibility
- No public API or wire-format removal: `message_end` gains an
additional optional `usage` field; existing consumers are unaffected.
- When a provider does not return authoritative usage, behavior falls
back to the previous token estimate (total only, no split).
- Non-agent flows (Dataflow `Pipeline`, sync `Graph.run`) are untouched.
## Testing
- [x] Simple agent answer: `message_end.usage.total_tokens` matches
provider usage.
- [x] Agent with cross-language retrieval: aggregate equals the sum of
both provider calls.
- [x] Tool-calling agent (multi-round): total accumulates across rounds.
- [x] Nested agent (agent-as-tool): sub-agent tokens included in the
parent run total.
- [x] Langfuse: agent generations show input/output split and are
grouped by session/user.
---------
Co-authored-by: yzc <yuzhichang@gmail.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
@@ -34,10 +34,12 @@ from peewee import fn
|
||||
class CanvasTemplateService(CommonService):
|
||||
model = CanvasTemplate
|
||||
|
||||
|
||||
class DataFlowTemplateService(CommonService):
|
||||
"""
|
||||
Alias of CanvasTemplateService
|
||||
"""
|
||||
|
||||
model = CanvasTemplate
|
||||
|
||||
|
||||
@@ -46,8 +48,7 @@ class UserCanvasService(CommonService):
|
||||
|
||||
@classmethod
|
||||
@DB.connection_context()
|
||||
def get_list(cls, tenant_id,
|
||||
page_number, items_per_page, orderby, desc, id, title, canvas_category=CanvasCategory.Agent):
|
||||
def get_list(cls, tenant_id, page_number, items_per_page, orderby, desc, id, title, canvas_category=CanvasCategory.Agent):
|
||||
agents = cls.model.select()
|
||||
if id:
|
||||
agents = agents.where(cls.model.id == id)
|
||||
@@ -68,20 +69,9 @@ class UserCanvasService(CommonService):
|
||||
@DB.connection_context()
|
||||
def get_all_agents_by_tenant_ids(cls, tenant_ids, user_id):
|
||||
# will get all permitted agents, be cautious
|
||||
fields = [
|
||||
cls.model.id,
|
||||
cls.model.avatar,
|
||||
cls.model.title,
|
||||
cls.model.permission,
|
||||
cls.model.canvas_type,
|
||||
cls.model.canvas_category
|
||||
]
|
||||
fields = [cls.model.id, cls.model.avatar, cls.model.title, cls.model.permission, cls.model.canvas_type, cls.model.canvas_category]
|
||||
# find team agents and owned agents
|
||||
agents = cls.model.select(*fields).where(
|
||||
(cls.model.user_id.in_(tenant_ids) & (cls.model.permission == TenantPermission.TEAM.value)) | (
|
||||
cls.model.user_id == user_id
|
||||
)
|
||||
)
|
||||
agents = cls.model.select(*fields).where((cls.model.user_id.in_(tenant_ids) & (cls.model.permission == TenantPermission.TEAM.value)) | (cls.model.user_id == user_id))
|
||||
# sort by create_time, asc
|
||||
agents.order_by(cls.model.create_time.asc())
|
||||
# maybe cause slow query by deep paginate, optimize later
|
||||
@@ -100,7 +90,6 @@ class UserCanvasService(CommonService):
|
||||
@DB.connection_context()
|
||||
def get_by_canvas_id(cls, pid):
|
||||
try:
|
||||
|
||||
fields = [
|
||||
cls.model.id,
|
||||
cls.model.avatar,
|
||||
@@ -115,11 +104,9 @@ class UserCanvasService(CommonService):
|
||||
cls.model.update_date,
|
||||
cls.model.canvas_category,
|
||||
User.nickname,
|
||||
User.avatar.alias('tenant_avatar'),
|
||||
User.avatar.alias("tenant_avatar"),
|
||||
]
|
||||
agents = cls.model.select(*fields) \
|
||||
.join(User, on=(cls.model.user_id == User.id)) \
|
||||
.where(cls.model.id == pid)
|
||||
agents = cls.model.select(*fields).join(User, on=(cls.model.user_id == User.id)).where(cls.model.id == pid)
|
||||
# obj = cls.model.query(id=pid)[0]
|
||||
return True, agents.dicts()[0]
|
||||
except Exception as e:
|
||||
@@ -129,14 +116,7 @@ class UserCanvasService(CommonService):
|
||||
@classmethod
|
||||
@DB.connection_context()
|
||||
def get_basic_info_by_canvas_ids(cls, canvas_id):
|
||||
fields = [
|
||||
cls.model.id,
|
||||
cls.model.avatar,
|
||||
cls.model.user_id,
|
||||
cls.model.title,
|
||||
cls.model.permission,
|
||||
cls.model.canvas_category
|
||||
]
|
||||
fields = [cls.model.id, cls.model.avatar, cls.model.user_id, cls.model.title, cls.model.permission, cls.model.canvas_category]
|
||||
return cls.model.select(*fields).where(cls.model.id.in_(canvas_id)).dicts()
|
||||
|
||||
@classmethod
|
||||
@@ -162,20 +142,26 @@ class UserCanvasService(CommonService):
|
||||
cls.model.permission,
|
||||
cls.model.user_id.alias("tenant_id"),
|
||||
User.nickname,
|
||||
User.avatar.alias('tenant_avatar'),
|
||||
User.avatar.alias("tenant_avatar"),
|
||||
cls.model.update_time,
|
||||
cls.model.canvas_type,
|
||||
cls.model.canvas_category,
|
||||
cls.model.tags,
|
||||
]
|
||||
if keywords:
|
||||
agents = cls.model.select(*fields).join(User, on=(cls.model.user_id == User.id)).where(
|
||||
(((cls.model.user_id.in_(joined_tenant_ids)) & (cls.model.permission == TenantPermission.TEAM.value)) | (cls.model.user_id == user_id)),
|
||||
(fn.LOWER(cls.model.title).contains(keywords.lower()))
|
||||
agents = (
|
||||
cls.model.select(*fields)
|
||||
.join(User, on=(cls.model.user_id == User.id))
|
||||
.where(
|
||||
(((cls.model.user_id.in_(joined_tenant_ids)) & (cls.model.permission == TenantPermission.TEAM.value)) | (cls.model.user_id == user_id)),
|
||||
(fn.LOWER(cls.model.title).contains(keywords.lower())),
|
||||
)
|
||||
)
|
||||
else:
|
||||
agents = cls.model.select(*fields).join(User, on=(cls.model.user_id == User.id)).where(
|
||||
(((cls.model.user_id.in_(joined_tenant_ids)) & (cls.model.permission == TenantPermission.TEAM.value)) | (cls.model.user_id == user_id))
|
||||
agents = (
|
||||
cls.model.select(*fields)
|
||||
.join(User, on=(cls.model.user_id == User.id))
|
||||
.where((((cls.model.user_id.in_(joined_tenant_ids)) & (cls.model.permission == TenantPermission.TEAM.value)) | (cls.model.user_id == user_id)))
|
||||
)
|
||||
if canvas_category:
|
||||
agents = agents.where(cls.model.canvas_category == canvas_category)
|
||||
@@ -201,7 +187,7 @@ class UserCanvasService(CommonService):
|
||||
|
||||
# Get latest release time for each canvas
|
||||
if agents_list:
|
||||
canvas_ids = [a['id'] for a in agents_list]
|
||||
canvas_ids = [a["id"] for a in agents_list]
|
||||
release_times = (
|
||||
UserCanvasVersion.select(UserCanvasVersion.user_canvas_id, fn.MAX(UserCanvasVersion.create_time).alias("release_time"))
|
||||
.where((UserCanvasVersion.user_canvas_id.in_(canvas_ids)) & (UserCanvasVersion.release))
|
||||
@@ -210,7 +196,7 @@ class UserCanvasService(CommonService):
|
||||
release_time_map = {r.user_canvas_id: r.release_time for r in release_times}
|
||||
|
||||
for agent in agents_list:
|
||||
agent['release_time'] = release_time_map.get(agent['id'])
|
||||
agent["release_time"] = release_time_map.get(agent["id"])
|
||||
|
||||
return agents_list, count
|
||||
|
||||
@@ -218,9 +204,7 @@ class UserCanvasService(CommonService):
|
||||
@DB.connection_context()
|
||||
def list_tags(cls, joined_tenant_ids, user_id, canvas_category=None):
|
||||
"""Return {tag: agent_count} aggregated across agents visible to the user."""
|
||||
query = cls.model.select(cls.model.tags).where(
|
||||
((cls.model.user_id.in_(joined_tenant_ids)) & (cls.model.permission == TenantPermission.TEAM.value)) | (cls.model.user_id == user_id)
|
||||
)
|
||||
query = cls.model.select(cls.model.tags).where(((cls.model.user_id.in_(joined_tenant_ids)) & (cls.model.permission == TenantPermission.TEAM.value)) | (cls.model.user_id == user_id))
|
||||
if canvas_category:
|
||||
query = query.where(cls.model.canvas_category == canvas_category)
|
||||
|
||||
@@ -281,6 +265,7 @@ class UserCanvasService(CommonService):
|
||||
@DB.connection_context()
|
||||
def accessible(cls, canvas_id, tenant_id):
|
||||
from api.db.services.user_service import UserTenantService
|
||||
|
||||
e, c = UserCanvasService.get_by_canvas_id(canvas_id)
|
||||
if not e:
|
||||
return False
|
||||
@@ -345,18 +330,15 @@ async def completion(tenant_id, agent_id, session_id=None, **kwargs):
|
||||
conv = API4Conversation(**conv)
|
||||
|
||||
message_id = str(uuid4())
|
||||
conv.message.append({
|
||||
"role": "user",
|
||||
"content": query,
|
||||
"id": message_id,
|
||||
"files": files
|
||||
})
|
||||
conv.message.append({"role": "user", "content": query, "id": message_id, "files": files})
|
||||
txt = ""
|
||||
run_kwargs = {
|
||||
"query": query,
|
||||
"files": files,
|
||||
"user_id": user_id,
|
||||
"inputs": inputs,
|
||||
# Used by Canvas.run to correlate RAGFlow's Langfuse generations by session.
|
||||
"session_id": session_id,
|
||||
}
|
||||
if chat_template_kwargs is not None:
|
||||
run_kwargs["chat_template_kwargs"] = chat_template_kwargs
|
||||
@@ -394,14 +376,7 @@ async def completion_openai(tenant_id, agent_id, question, session_id=None, stre
|
||||
if stream:
|
||||
completion_tokens = 0
|
||||
try:
|
||||
async for ans in completion(
|
||||
tenant_id=tenant_id,
|
||||
agent_id=agent_id,
|
||||
session_id=session_id,
|
||||
query=question,
|
||||
user_id=user_id,
|
||||
**kwargs
|
||||
):
|
||||
async for ans in completion(tenant_id=tenant_id, agent_id=agent_id, session_id=session_id, query=question, user_id=user_id, **kwargs):
|
||||
if isinstance(ans, str):
|
||||
try:
|
||||
ans = json.loads(ans[5:]) # remove "data:"
|
||||
@@ -417,14 +392,7 @@ async def completion_openai(tenant_id, agent_id, question, session_id=None, stre
|
||||
|
||||
completion_tokens += len(tiktoken_encoder.encode(content_piece))
|
||||
|
||||
openai_data = get_data_openai(
|
||||
id=session_id or str(uuid4()),
|
||||
model=agent_id,
|
||||
content=content_piece,
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
stream=True
|
||||
)
|
||||
openai_data = get_data_openai(id=session_id or str(uuid4()), model=agent_id, content=content_piece, prompt_tokens=prompt_tokens, completion_tokens=completion_tokens, stream=True)
|
||||
|
||||
if ans.get("data", {}).get("reference", None):
|
||||
openai_data["choices"][0]["delta"]["reference"] = ans["data"]["reference"]
|
||||
@@ -435,32 +403,29 @@ async def completion_openai(tenant_id, agent_id, question, session_id=None, stre
|
||||
|
||||
except Exception as e:
|
||||
logging.exception(e)
|
||||
yield "data: " + json.dumps(
|
||||
get_data_openai(
|
||||
id=session_id or str(uuid4()),
|
||||
model=agent_id,
|
||||
content=f"**ERROR**: {str(e)}",
|
||||
finish_reason="stop",
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=len(tiktoken_encoder.encode(f"**ERROR**: {str(e)}")),
|
||||
stream=True
|
||||
),
|
||||
ensure_ascii=False
|
||||
) + "\n\n"
|
||||
yield (
|
||||
"data: "
|
||||
+ json.dumps(
|
||||
get_data_openai(
|
||||
id=session_id or str(uuid4()),
|
||||
model=agent_id,
|
||||
content=f"**ERROR**: {str(e)}",
|
||||
finish_reason="stop",
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=len(tiktoken_encoder.encode(f"**ERROR**: {str(e)}")),
|
||||
stream=True,
|
||||
),
|
||||
ensure_ascii=False,
|
||||
)
|
||||
+ "\n\n"
|
||||
)
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
else:
|
||||
try:
|
||||
all_content = ""
|
||||
reference = {}
|
||||
async for ans in completion(
|
||||
tenant_id=tenant_id,
|
||||
agent_id=agent_id,
|
||||
session_id=session_id,
|
||||
query=question,
|
||||
user_id=user_id,
|
||||
**kwargs
|
||||
):
|
||||
async for ans in completion(tenant_id=tenant_id, agent_id=agent_id, session_id=session_id, query=question, user_id=user_id, **kwargs):
|
||||
if isinstance(ans, str):
|
||||
ans = json.loads(ans[5:])
|
||||
if ans.get("event") not in ["message", "message_end"]:
|
||||
@@ -475,13 +440,7 @@ async def completion_openai(tenant_id, agent_id, question, session_id=None, stre
|
||||
completion_tokens = len(tiktoken_encoder.encode(all_content))
|
||||
|
||||
openai_data = get_data_openai(
|
||||
id=session_id or str(uuid4()),
|
||||
model=agent_id,
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
content=all_content,
|
||||
finish_reason="stop",
|
||||
param=None
|
||||
id=session_id or str(uuid4()), model=agent_id, prompt_tokens=prompt_tokens, completion_tokens=completion_tokens, content=all_content, finish_reason="stop", param=None
|
||||
)
|
||||
|
||||
if reference:
|
||||
@@ -497,5 +456,5 @@ async def completion_openai(tenant_id, agent_id, question, session_id=None, stre
|
||||
completion_tokens=len(tiktoken_encoder.encode(f"**ERROR**: {str(e)}")),
|
||||
content=f"**ERROR**: {str(e)}",
|
||||
finish_reason="stop",
|
||||
param=None
|
||||
param=None,
|
||||
)
|
||||
|
||||
@@ -27,7 +27,7 @@ from langfuse import propagate_attributes
|
||||
from api.db.db_models import LLM
|
||||
from api.db.services.common_service import CommonService
|
||||
from api.db.services.tenant_llm_service import LLM4Tenant
|
||||
from common.token_utils import num_tokens_from_string
|
||||
from common.token_utils import num_tokens_from_string, record_run_token_usage, langfuse_run_attrs
|
||||
|
||||
|
||||
class LLMService(CommonService):
|
||||
@@ -39,11 +39,49 @@ class LLMBundle(LLM4Tenant):
|
||||
super().__init__(tenant_id, model_config, lang, **kwargs)
|
||||
|
||||
def _start_langfuse_observation(self, **kwargs):
|
||||
# Correlating attributes (session_id/user_id) let Langfuse group all of a
|
||||
# turn's generations. They may come from this bundle (chat/dialog path) or,
|
||||
# for agent runs whose bundles are created without them, from the per-run
|
||||
# context installed by Canvas.run.
|
||||
attrs = {}
|
||||
if self.langfuse_session_id:
|
||||
with propagate_attributes(session_id=self.langfuse_session_id):
|
||||
attrs["session_id"] = self.langfuse_session_id
|
||||
run_attrs = langfuse_run_attrs.get()
|
||||
if run_attrs:
|
||||
for k in ("session_id", "user_id"):
|
||||
if run_attrs.get(k) and k not in attrs:
|
||||
attrs[k] = run_attrs[k]
|
||||
if attrs:
|
||||
with propagate_attributes(**attrs):
|
||||
return self.langfuse.start_observation(**kwargs)
|
||||
return self.langfuse.start_observation(**kwargs)
|
||||
|
||||
def _reset_last_usage(self) -> None:
|
||||
"""Clear the model's per-call usage so a failed call that returns before
|
||||
updating it cannot leak the previous call's usage into this run."""
|
||||
if hasattr(self.mdl, "last_usage"):
|
||||
self.mdl.last_usage = {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
|
||||
|
||||
def _report_usage(self, total_tokens: int) -> dict:
|
||||
"""Record a chat call's usage to the active agent run and return the
|
||||
prompt/completion/total split for Langfuse.
|
||||
|
||||
``total_tokens`` is the authoritative total from the call. The prompt/completion
|
||||
split is taken from the provider response (``mdl.last_usage``) only when it is
|
||||
consistent with ``total_tokens`` (i.e. produced by this same call); otherwise the
|
||||
split is reported as 0 while the total still aggregates correctly.
|
||||
"""
|
||||
split = getattr(self.mdl, "last_usage", None) or {}
|
||||
prompt = int(split.get("prompt_tokens", 0) or 0)
|
||||
completion = int(split.get("completion_tokens", 0) or 0)
|
||||
if not total_tokens:
|
||||
total_tokens = int(split.get("total_tokens", 0) or 0)
|
||||
if (prompt + completion) != total_tokens:
|
||||
# Stale or inconsistent split — keep the total, drop the unreliable split.
|
||||
prompt, completion = 0, 0
|
||||
record_run_token_usage(prompt, completion, total_tokens)
|
||||
return {"input": prompt, "output": completion, "total": total_tokens}
|
||||
|
||||
def close(self):
|
||||
"""Release resources held by this LLMBundle instance."""
|
||||
super().close()
|
||||
@@ -139,7 +177,9 @@ class LLMBundle(LLM4Tenant):
|
||||
|
||||
def similarity(self, query: str, texts: list):
|
||||
if self.langfuse:
|
||||
generation = self._start_langfuse_observation(trace_context=self.trace_context, as_type="generation", name="similarity", model=self.model_config["llm_name"], input={"query": query, "texts": texts})
|
||||
generation = self._start_langfuse_observation(
|
||||
trace_context=self.trace_context, as_type="generation", name="similarity", model=self.model_config["llm_name"], input={"query": query, "texts": texts}
|
||||
)
|
||||
|
||||
sim, used_tokens = self.mdl.similarity(query, texts)
|
||||
logging.info("LLMBundle.similarity used_tokens: %d", used_tokens)
|
||||
@@ -165,7 +205,9 @@ class LLMBundle(LLM4Tenant):
|
||||
|
||||
def describe_with_prompt(self, image, prompt):
|
||||
if self.langfuse:
|
||||
generation = self._start_langfuse_observation(trace_context=self.trace_context, as_type="generation", name="describe_with_prompt", metadata={"model": self.model_config["llm_name"], "prompt": prompt})
|
||||
generation = self._start_langfuse_observation(
|
||||
trace_context=self.trace_context, as_type="generation", name="describe_with_prompt", metadata={"model": self.model_config["llm_name"], "prompt": prompt}
|
||||
)
|
||||
|
||||
txt, used_tokens = self.mdl.describe_with_prompt(image, prompt)
|
||||
logging.info("LLMBundle.describe_with_prompt used_tokens: %d", used_tokens)
|
||||
@@ -194,7 +236,8 @@ class LLMBundle(LLM4Tenant):
|
||||
supports_stream = hasattr(mdl, "stream_transcription") and callable(getattr(mdl, "stream_transcription"))
|
||||
if supports_stream:
|
||||
if self.langfuse:
|
||||
generation = self._start_langfuse_observation(as_type="generation",
|
||||
generation = self._start_langfuse_observation(
|
||||
as_type="generation",
|
||||
trace_context=self.trace_context,
|
||||
name="stream_transcription",
|
||||
metadata={"model": self.model_config["llm_name"]},
|
||||
@@ -228,7 +271,8 @@ class LLMBundle(LLM4Tenant):
|
||||
return
|
||||
|
||||
if self.langfuse:
|
||||
generation = self._start_langfuse_observation(as_type="generation",
|
||||
generation = self._start_langfuse_observation(
|
||||
as_type="generation",
|
||||
trace_context=self.trace_context,
|
||||
name="stream_transcription",
|
||||
metadata={"model": self.model_config["llm_name"]},
|
||||
@@ -377,11 +421,14 @@ class LLMBundle(LLM4Tenant):
|
||||
|
||||
generation = None
|
||||
if self.langfuse:
|
||||
generation = self._start_langfuse_observation(trace_context=self.trace_context, as_type="generation", name="chat", model=self.model_config["llm_name"], input={"system": system, "history": history})
|
||||
generation = self._start_langfuse_observation(
|
||||
trace_context=self.trace_context, as_type="generation", name="chat", model=self.model_config["llm_name"], input={"system": system, "history": history}
|
||||
)
|
||||
|
||||
chat_partial = partial(base_fn, system, history, gen_conf)
|
||||
use_kwargs = self._clean_param(chat_partial, **kwargs)
|
||||
|
||||
self._reset_last_usage()
|
||||
try:
|
||||
txt, used_tokens = await chat_partial(**use_kwargs)
|
||||
except Exception as e:
|
||||
@@ -397,8 +444,10 @@ class LLMBundle(LLM4Tenant):
|
||||
if used_tokens:
|
||||
logging.info("LLMBundle.async_chat used_tokens: %d", used_tokens)
|
||||
|
||||
usage_details = self._report_usage(used_tokens)
|
||||
|
||||
if generation:
|
||||
generation.update(output={"output": txt}, usage_details={"total_tokens": used_tokens})
|
||||
generation.update(output={"output": txt}, usage_details=usage_details)
|
||||
generation.end()
|
||||
|
||||
return txt
|
||||
@@ -418,11 +467,14 @@ class LLMBundle(LLM4Tenant):
|
||||
|
||||
generation = None
|
||||
if self.langfuse:
|
||||
generation = self._start_langfuse_observation(trace_context=self.trace_context, as_type="generation", name="chat_streamly", model=self.model_config["llm_name"], input={"system": system, "history": history})
|
||||
generation = self._start_langfuse_observation(
|
||||
trace_context=self.trace_context, as_type="generation", name="chat_streamly", model=self.model_config["llm_name"], input={"system": system, "history": history}
|
||||
)
|
||||
|
||||
if stream_fn:
|
||||
chat_partial = partial(stream_fn, system, history, gen_conf)
|
||||
use_kwargs = self._clean_param(chat_partial, **kwargs)
|
||||
self._reset_last_usage()
|
||||
try:
|
||||
async for txt in chat_partial(**use_kwargs):
|
||||
if isinstance(txt, int):
|
||||
@@ -444,8 +496,9 @@ class LLMBundle(LLM4Tenant):
|
||||
raise
|
||||
if total_tokens:
|
||||
logging.info("LLMBundle.async_chat_streamly used_tokens: %d", total_tokens)
|
||||
usage_details = self._report_usage(total_tokens)
|
||||
if generation:
|
||||
generation.update(output={"output": ans}, usage_details={"total_tokens": total_tokens})
|
||||
generation.update(output={"output": ans}, usage_details=usage_details)
|
||||
generation.end()
|
||||
return
|
||||
|
||||
@@ -461,11 +514,14 @@ class LLMBundle(LLM4Tenant):
|
||||
|
||||
generation = None
|
||||
if self.langfuse:
|
||||
generation = self._start_langfuse_observation(trace_context=self.trace_context, as_type="generation", name="chat_streamly", model=self.model_config["llm_name"], input={"system": system, "history": history})
|
||||
generation = self._start_langfuse_observation(
|
||||
trace_context=self.trace_context, as_type="generation", name="chat_streamly", model=self.model_config["llm_name"], input={"system": system, "history": history}
|
||||
)
|
||||
|
||||
if stream_fn:
|
||||
chat_partial = partial(stream_fn, system, history, gen_conf)
|
||||
use_kwargs = self._clean_param(chat_partial, **kwargs)
|
||||
self._reset_last_usage()
|
||||
try:
|
||||
async for txt in chat_partial(**use_kwargs):
|
||||
if isinstance(txt, int):
|
||||
@@ -487,7 +543,8 @@ class LLMBundle(LLM4Tenant):
|
||||
raise
|
||||
if total_tokens:
|
||||
logging.info("LLMBundle.async_chat_streamly_delta used_tokens: %d", total_tokens)
|
||||
usage_details = self._report_usage(total_tokens)
|
||||
if generation:
|
||||
generation.update(output={"output": ans}, usage_details={"total_tokens": total_tokens})
|
||||
generation.update(output={"output": ans}, usage_details=usage_details)
|
||||
generation.end()
|
||||
return
|
||||
|
||||
Reference in New Issue
Block a user