mirror of
https://github.com/langgenius/dify.git
synced 2026-09-05 00:31:19 +08:00
fix(agent): retain usage for incomplete runs (#40937)
This commit is contained in:
parent
2400c6c59d
commit
c38fb16efa
@ -100,6 +100,7 @@ class AgentBackendRunFailedInternalEvent(AgentBackendInternalEventBase):
|
||||
error_type: RunFailureType | None = None
|
||||
reason: str | None = None
|
||||
session_snapshot: CompositorSessionSnapshot | None = None
|
||||
usage: dict[str, JsonValue] | None = None
|
||||
|
||||
|
||||
class AgentBackendRunCancelledInternalEvent(AgentBackendInternalEventBase):
|
||||
@ -109,6 +110,7 @@ class AgentBackendRunCancelledInternalEvent(AgentBackendInternalEventBase):
|
||||
reason: str | None = None
|
||||
message: str | None = None
|
||||
session_snapshot: CompositorSessionSnapshot | None = None
|
||||
usage: dict[str, JsonValue] | None = None
|
||||
|
||||
|
||||
type AgentBackendInternalEvent = Annotated[
|
||||
@ -187,6 +189,7 @@ class AgentBackendRunEventAdapter:
|
||||
error_type=event.data.error_type,
|
||||
reason=event.data.reason,
|
||||
session_snapshot=event.data.session_snapshot,
|
||||
usage=_agent_run_usage(event.data.usage),
|
||||
)
|
||||
]
|
||||
case RunCancelledEvent():
|
||||
@ -197,6 +200,7 @@ class AgentBackendRunEventAdapter:
|
||||
reason=event.data.reason,
|
||||
message=event.data.message,
|
||||
session_snapshot=event.data.session_snapshot,
|
||||
usage=_agent_run_usage(event.data.usage),
|
||||
)
|
||||
]
|
||||
raise TypeError(f"unsupported agent backend run event: {type(event).__name__}")
|
||||
|
||||
@ -68,7 +68,7 @@ from graphon.model_runtime.errors.invoke import (
|
||||
from models.agent import AgentConfigVersionKind
|
||||
from models.agent_config_entities import AgentSoulConfig
|
||||
from models.enums import CreatorUserRole
|
||||
from models.model import MessageAgentThought
|
||||
from models.model import Message, MessageAgentThought
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@ -703,6 +703,19 @@ class AgentAppRunner:
|
||||
)
|
||||
return
|
||||
|
||||
terminal_usage = None
|
||||
if isinstance(
|
||||
terminal,
|
||||
AgentBackendRunSucceededInternalEvent
|
||||
| AgentBackendRunFailedInternalEvent
|
||||
| AgentBackendRunCancelledInternalEvent,
|
||||
):
|
||||
terminal_usage = _llm_usage_from_agent_backend(terminal.usage)
|
||||
self._persist_message_usage(
|
||||
message_id=message_id,
|
||||
usage=terminal_usage,
|
||||
)
|
||||
|
||||
if isinstance(terminal, AgentBackendRunFailedInternalEvent | AgentBackendRunCancelledInternalEvent):
|
||||
# None means no post-exit snapshot was produced; leave the previously stored session snapshot untouched.
|
||||
if terminal.session_snapshot is not None:
|
||||
@ -736,7 +749,7 @@ class AgentAppRunner:
|
||||
model_name=model_name,
|
||||
answer=answer,
|
||||
query=query,
|
||||
usage=_llm_usage_from_agent_backend(terminal.usage),
|
||||
usage=terminal_usage,
|
||||
)
|
||||
self._save_session(
|
||||
scope=scope,
|
||||
@ -856,6 +869,7 @@ class AgentAppRunner:
|
||||
model_name=model_name,
|
||||
answer=self._ask_human_message(created.args),
|
||||
query=query,
|
||||
usage=_llm_usage_from_agent_backend(terminal.usage),
|
||||
)
|
||||
|
||||
def _resolve_pending_ask_human(
|
||||
@ -958,6 +972,7 @@ class AgentAppRunner:
|
||||
after=last_event_id,
|
||||
session_scope=session_scope,
|
||||
binding_id=binding_id,
|
||||
message_id=message_id,
|
||||
)
|
||||
raise GenerateTaskStoppedError()
|
||||
for internal_event in self._event_adapter.adapt(public_event):
|
||||
@ -968,6 +983,7 @@ class AgentAppRunner:
|
||||
after=last_event_id,
|
||||
session_scope=session_scope,
|
||||
binding_id=binding_id,
|
||||
message_id=message_id,
|
||||
)
|
||||
raise GenerateTaskStoppedError()
|
||||
if internal_event.type in (
|
||||
@ -1010,6 +1026,7 @@ class AgentAppRunner:
|
||||
after=last_event_id,
|
||||
session_scope=session_scope,
|
||||
binding_id=binding_id,
|
||||
message_id=message_id,
|
||||
)
|
||||
if queue_manager.is_stopped():
|
||||
raise GenerateTaskStoppedError() from error
|
||||
@ -1021,6 +1038,7 @@ class AgentAppRunner:
|
||||
after=last_event_id,
|
||||
session_scope=session_scope,
|
||||
binding_id=binding_id,
|
||||
message_id=message_id,
|
||||
)
|
||||
raise GenerateTaskStoppedError()
|
||||
return terminal, process_recorder
|
||||
@ -1032,6 +1050,7 @@ class AgentAppRunner:
|
||||
after: str | None,
|
||||
session_scope: AgentAppSessionScope,
|
||||
binding_id: str,
|
||||
message_id: str,
|
||||
) -> None:
|
||||
try:
|
||||
public_event = self._agent_backend_client.cancel_run_and_wait(run_id, after=after)
|
||||
@ -1045,6 +1064,11 @@ class AgentAppRunner:
|
||||
binding_id=binding_id,
|
||||
snapshot=internal_event.session_snapshot,
|
||||
)
|
||||
if isinstance(internal_event, AgentBackendRunCancelledInternalEvent):
|
||||
self._persist_message_usage(
|
||||
message_id=message_id,
|
||||
usage=_llm_usage_from_agent_backend(internal_event.usage),
|
||||
)
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"Failed to finish cancelling stopped Agent App backend run: run_id=%s",
|
||||
@ -1053,11 +1077,23 @@ class AgentAppRunner:
|
||||
)
|
||||
|
||||
def _publish_answer(
|
||||
self, *, queue_manager: AppQueueManager, model_name: str, answer: str, query: str | None
|
||||
self,
|
||||
*,
|
||||
queue_manager: AppQueueManager,
|
||||
model_name: str,
|
||||
answer: str,
|
||||
query: str | None,
|
||||
usage: LLMUsage | None = None,
|
||||
) -> None:
|
||||
# MVP: emit the full answer as a single chunk + message-end. The chat
|
||||
# task pipeline streams the chunk over SSE and persists the message.
|
||||
publish_text_answer(queue_manager=queue_manager, model_name=model_name, answer=answer, user_query=query)
|
||||
publish_text_answer(
|
||||
queue_manager=queue_manager,
|
||||
model_name=model_name,
|
||||
answer=answer,
|
||||
user_query=query,
|
||||
usage=usage,
|
||||
)
|
||||
|
||||
def _publish_terminal_answer(
|
||||
self,
|
||||
@ -1077,6 +1113,40 @@ class AgentAppRunner:
|
||||
user_query=query,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _persist_message_usage(*, message_id: str, usage: LLMUsage | None) -> None:
|
||||
"""Persist terminal usage independently of the client-facing stream lifecycle."""
|
||||
if usage is None or (usage.total_tokens <= 0 and usage.total_price <= 0):
|
||||
return
|
||||
try:
|
||||
message = db.session.get(Message, message_id)
|
||||
if message is None:
|
||||
logger.warning("Cannot persist Agent App usage: message not found: %s", message_id)
|
||||
return
|
||||
message.message_tokens = usage.prompt_tokens
|
||||
message.message_unit_price = usage.prompt_unit_price
|
||||
message.message_price_unit = usage.prompt_price_unit
|
||||
message.answer_tokens = usage.completion_tokens
|
||||
message.answer_unit_price = usage.completion_unit_price
|
||||
message.answer_price_unit = usage.completion_price_unit
|
||||
message.total_price = usage.total_price
|
||||
message.currency = usage.currency
|
||||
if usage.latency > 0:
|
||||
message.provider_response_latency = usage.latency
|
||||
|
||||
try:
|
||||
metadata = json.loads(message.message_metadata) if message.message_metadata else {}
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
metadata = {}
|
||||
if not isinstance(metadata, dict):
|
||||
metadata = {}
|
||||
metadata["usage"] = usage.model_dump(mode="json")
|
||||
message.message_metadata = json.dumps(metadata, ensure_ascii=False)
|
||||
db.session.commit()
|
||||
except Exception:
|
||||
db.session.rollback()
|
||||
logger.warning("Failed to persist partial Agent App usage: message_id=%s", message_id, exc_info=True)
|
||||
|
||||
def _save_session(
|
||||
self,
|
||||
*,
|
||||
|
||||
@ -1,3 +1,4 @@
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from collections.abc import Generator, Mapping, Sequence
|
||||
@ -292,8 +293,16 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline[EasyUIAppGenerat
|
||||
)
|
||||
|
||||
with session_factory.create_session() as session:
|
||||
# Save message
|
||||
self._save_message(session=session, trace_manager=trace_manager)
|
||||
# A stopped Agent run may persist provider-reported usage after
|
||||
# cancellation completes. Do not replace it with local token estimates.
|
||||
if isinstance(event, QueueStopEvent):
|
||||
self._save_message(
|
||||
session=session,
|
||||
trace_manager=trace_manager,
|
||||
preserve_existing_usage=True,
|
||||
)
|
||||
else:
|
||||
self._save_message(session=session, trace_manager=trace_manager)
|
||||
session.commit()
|
||||
message_end_resp = self._message_end_to_stream_response()
|
||||
yield message_end_resp
|
||||
@ -389,7 +398,13 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline[EasyUIAppGenerat
|
||||
continue
|
||||
return delta_text
|
||||
|
||||
def _save_message(self, *, session: Session, trace_manager: TraceQueueManager | None = None):
|
||||
def _save_message(
|
||||
self,
|
||||
*,
|
||||
session: Session,
|
||||
trace_manager: TraceQueueManager | None = None,
|
||||
preserve_existing_usage: bool = False,
|
||||
):
|
||||
"""
|
||||
Save message.
|
||||
:return:
|
||||
@ -410,24 +425,39 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline[EasyUIAppGenerat
|
||||
self._model_config.mode, self._task_state.llm_result.prompt_messages
|
||||
)
|
||||
object.__setattr__(message, "message", saved_prompt)
|
||||
message.message_tokens = usage.prompt_tokens
|
||||
message.message_unit_price = usage.prompt_unit_price
|
||||
message.message_price_unit = usage.prompt_price_unit
|
||||
try:
|
||||
existing_metadata = json.loads(message.message_metadata) if message.message_metadata else {}
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
existing_metadata = {}
|
||||
if not isinstance(existing_metadata, dict):
|
||||
existing_metadata = {}
|
||||
has_persisted_usage = preserve_existing_usage and (
|
||||
int(message.message_tokens or 0) + int(message.answer_tokens or 0) > 0 or bool(message.total_price)
|
||||
)
|
||||
if not has_persisted_usage:
|
||||
message.message_tokens = usage.prompt_tokens
|
||||
message.message_unit_price = usage.prompt_unit_price
|
||||
message.message_price_unit = usage.prompt_price_unit
|
||||
message.answer = (
|
||||
PromptTemplateParser.remove_template_variables(llm_result.message.get_text_content().strip())
|
||||
if llm_result.message.content
|
||||
else ""
|
||||
)
|
||||
message.updated_at = naive_utc_now()
|
||||
message.answer_tokens = usage.completion_tokens
|
||||
message.answer_unit_price = usage.completion_unit_price
|
||||
message.answer_price_unit = usage.completion_price_unit
|
||||
message.provider_response_latency = time.perf_counter() - self.start_at
|
||||
message.total_price = usage.total_price
|
||||
message.currency = usage.currency
|
||||
self._task_state.llm_result.usage.latency = message.provider_response_latency
|
||||
self._task_state.metadata.usage = self._task_state.llm_result.usage
|
||||
message.message_metadata = self._task_state.metadata.model_dump_json()
|
||||
if not has_persisted_usage:
|
||||
message.answer_tokens = usage.completion_tokens
|
||||
message.answer_unit_price = usage.completion_unit_price
|
||||
message.answer_price_unit = usage.completion_price_unit
|
||||
message.provider_response_latency = time.perf_counter() - self.start_at
|
||||
message.total_price = usage.total_price
|
||||
message.currency = usage.currency
|
||||
self._task_state.llm_result.usage.latency = message.provider_response_latency
|
||||
self._task_state.metadata.usage = self._task_state.llm_result.usage
|
||||
|
||||
metadata = self._task_state.metadata.model_dump(mode="json")
|
||||
if has_persisted_usage and "usage" in existing_metadata:
|
||||
metadata["usage"] = existing_metadata["usage"]
|
||||
message.message_metadata = json.dumps(metadata, ensure_ascii=False)
|
||||
|
||||
if trace_manager:
|
||||
trace_manager.add_trace_task(
|
||||
|
||||
@ -331,7 +331,13 @@ class WorkflowAgentOutputAdapter:
|
||||
}
|
||||
)
|
||||
session_snapshot = None
|
||||
if isinstance(event, AgentBackendRunSucceededInternalEvent | AgentBackendDeferredToolCallInternalEvent):
|
||||
if isinstance(
|
||||
event,
|
||||
AgentBackendRunSucceededInternalEvent
|
||||
| AgentBackendDeferredToolCallInternalEvent
|
||||
| AgentBackendRunFailedInternalEvent
|
||||
| AgentBackendRunCancelledInternalEvent,
|
||||
):
|
||||
session_snapshot = event.session_snapshot
|
||||
if event.usage is not None:
|
||||
agent_backend["usage"] = dict(event.usage)
|
||||
|
||||
@ -138,6 +138,7 @@ def test_event_adapter_maps_run_failed_to_failed_result():
|
||||
error="boom",
|
||||
error_type=RunFailureType.AGENT_RUN_LIMIT_EXCEEDED,
|
||||
reason="runtime",
|
||||
usage=AgentRunUsage(prompt_tokens=13, completion_tokens=8),
|
||||
),
|
||||
)
|
||||
)
|
||||
@ -149,6 +150,22 @@ def test_event_adapter_maps_run_failed_to_failed_result():
|
||||
error="boom",
|
||||
error_type=RunFailureType.AGENT_RUN_LIMIT_EXCEEDED,
|
||||
reason="runtime",
|
||||
usage={
|
||||
"prompt_tokens": 13,
|
||||
"completion_tokens": 8,
|
||||
"total_tokens": 21,
|
||||
"prompt_unit_price": "0",
|
||||
"prompt_price_unit": "0",
|
||||
"prompt_price": "0",
|
||||
"completion_unit_price": "0",
|
||||
"completion_price_unit": "0",
|
||||
"completion_price": "0",
|
||||
"total_price": "0",
|
||||
"currency": "USD",
|
||||
"latency": 0.0,
|
||||
"time_to_first_token": None,
|
||||
"time_to_generate": None,
|
||||
},
|
||||
)
|
||||
]
|
||||
|
||||
@ -254,7 +271,11 @@ def test_event_adapter_maps_run_cancelled_to_terminal_cancelled():
|
||||
RunCancelledEvent(
|
||||
id="6-0",
|
||||
run_id="run-1",
|
||||
data=RunCancelledEventData(reason="user_cancelled", message="Stopped by user"),
|
||||
data=RunCancelledEventData(
|
||||
reason="user_cancelled",
|
||||
message="Stopped by user",
|
||||
usage=AgentRunUsage(prompt_tokens=5, completion_tokens=3),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
@ -264,5 +285,21 @@ def test_event_adapter_maps_run_cancelled_to_terminal_cancelled():
|
||||
source_event_id="6-0",
|
||||
reason="user_cancelled",
|
||||
message="Stopped by user",
|
||||
usage={
|
||||
"prompt_tokens": 5,
|
||||
"completion_tokens": 3,
|
||||
"total_tokens": 8,
|
||||
"prompt_unit_price": "0",
|
||||
"prompt_price_unit": "0",
|
||||
"prompt_price": "0",
|
||||
"completion_unit_price": "0",
|
||||
"completion_price_unit": "0",
|
||||
"completion_price": "0",
|
||||
"total_price": "0",
|
||||
"currency": "USD",
|
||||
"latency": 0.0,
|
||||
"time_to_first_token": None,
|
||||
"time_to_generate": None,
|
||||
},
|
||||
)
|
||||
]
|
||||
|
||||
@ -4,6 +4,7 @@ saved, using the deterministic fake backend client (no live stack)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Callable, Iterator
|
||||
from datetime import UTC, datetime
|
||||
from decimal import Decimal
|
||||
@ -63,10 +64,11 @@ from core.app.entities.queue_entities import (
|
||||
)
|
||||
from core.workflow.nodes.agent_v2.ask_human_resume import AskHumanResumeOutcome
|
||||
from core.workflow.nodes.agent_v2.dify_tools_builder import WorkflowAgentToolLayers
|
||||
from graphon.model_runtime.entities.llm_entities import LLMResult
|
||||
from graphon.model_runtime.entities.llm_entities import LLMResult, LLMUsage
|
||||
from graphon.model_runtime.errors.invoke import InvokeRateLimitError
|
||||
from models.agent_config_entities import AgentSoulConfig
|
||||
from models.model import MessageAgentThought
|
||||
from models.enums import ConversationFromSource
|
||||
from models.model import AppMode, Message, MessageAgentThought
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
@ -148,6 +150,72 @@ class _CancelAndWaitFailingClient(_RecordingFakeAgentBackendRunClient):
|
||||
raise RuntimeError(f"failed to finish cancelling {run_id}")
|
||||
|
||||
|
||||
class _UsageCancellationClient(_RecordingFakeAgentBackendRunClient):
|
||||
@override
|
||||
def cancel_run_and_wait(
|
||||
self,
|
||||
run_id: str,
|
||||
request: CancelRunRequest | None = None,
|
||||
*,
|
||||
after: str | None = None,
|
||||
) -> RunCancelledEvent:
|
||||
event = super().cancel_run_and_wait(run_id, request=request, after=after)
|
||||
return event.model_copy(
|
||||
update={
|
||||
"data": event.data.model_copy(
|
||||
update={"usage": AgentRunUsage(prompt_tokens=13, completion_tokens=8, total_price=Decimal("0.21"))}
|
||||
)
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class _UsageFailedClient(FakeAgentBackendRunClient):
|
||||
@override
|
||||
def stream_events(
|
||||
self,
|
||||
run_id: str,
|
||||
*,
|
||||
after: str | None = None,
|
||||
should_stop: Callable[[], bool] | None = None,
|
||||
) -> Iterator[RunEvent]:
|
||||
del after, should_stop
|
||||
yield RunStartedEvent(id="1-0", run_id=run_id)
|
||||
yield RunFailedEvent(
|
||||
id="2-0",
|
||||
run_id=run_id,
|
||||
data=RunFailedEventData(
|
||||
error="failed after model calls",
|
||||
usage=AgentRunUsage(
|
||||
prompt_tokens=7,
|
||||
completion_tokens=5,
|
||||
total_price=Decimal("0.12"),
|
||||
latency=0.4,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class _UsagePausedClient(FakeAgentBackendRunClient):
|
||||
def __init__(self) -> None:
|
||||
super().__init__(scenario=FakeAgentBackendScenario.PAUSED)
|
||||
|
||||
@override
|
||||
def _events(self, run_id: str) -> tuple[RunEvent, ...]:
|
||||
events = super()._events(run_id)
|
||||
terminal = events[-1]
|
||||
assert isinstance(terminal, RunSucceededEvent)
|
||||
return (
|
||||
*events[:-1],
|
||||
terminal.model_copy(
|
||||
update={
|
||||
"data": terminal.data.model_copy(
|
||||
update={"usage": AgentRunUsage(prompt_tokens=5, completion_tokens=3)}
|
||||
)
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class _RunLimitBindingLostFakeAgentBackendRunClient(FakeAgentBackendRunClient):
|
||||
@override
|
||||
def stream_events(
|
||||
@ -592,6 +660,33 @@ def _run(runner: AgentAppRunner, qm: _FakeQueueManager) -> None:
|
||||
)
|
||||
|
||||
|
||||
def _message_record() -> Message:
|
||||
message = Message(
|
||||
app_id="app-1",
|
||||
conversation_id="conv-1",
|
||||
inputs={},
|
||||
query="hello",
|
||||
message={},
|
||||
message_tokens=0,
|
||||
message_unit_price=0,
|
||||
message_price_unit=0,
|
||||
answer="",
|
||||
answer_tokens=0,
|
||||
answer_unit_price=0,
|
||||
answer_price_unit=0,
|
||||
provider_response_latency=0,
|
||||
total_price=0,
|
||||
currency="USD",
|
||||
invoke_from=InvokeFrom.WEB_APP,
|
||||
from_source=ConversationFromSource.API,
|
||||
from_end_user_id="user-1",
|
||||
from_account_id=None,
|
||||
app_mode=AppMode.AGENT,
|
||||
)
|
||||
message.id = "msg-1"
|
||||
return message
|
||||
|
||||
|
||||
def _message_end(qm: _FakeQueueManager) -> QueueMessageEndEvent:
|
||||
return next(e for e in qm.events if isinstance(e, QueueMessageEndEvent))
|
||||
|
||||
@ -683,6 +778,25 @@ def test_successful_turn_routes_stream_text_to_agent_message_and_uses_terminal_o
|
||||
assert store.saved
|
||||
|
||||
|
||||
def test_successful_turn_persists_usage_without_a_queue_consumer(sqlite_session: Session) -> None:
|
||||
sqlite_session.add(_message_record())
|
||||
sqlite_session.flush()
|
||||
|
||||
_run(
|
||||
_runner(_StreamingFakeAgentBackendRunClient(), _FakeSessionStore()),
|
||||
_FakeQueueManager(),
|
||||
)
|
||||
|
||||
sqlite_session.expire_all()
|
||||
message = sqlite_session.get(Message, "msg-1")
|
||||
assert message is not None
|
||||
assert message.message_tokens == 3
|
||||
assert message.answer_tokens == 5
|
||||
assert message.total_price == Decimal("0.000165")
|
||||
assert message.message_metadata is not None
|
||||
assert json.loads(message.message_metadata)["usage"]["total_tokens"] == 8
|
||||
|
||||
|
||||
def test_successful_turn_routes_single_agent_message_delta(sqlite_session: Session) -> None:
|
||||
client = _StreamingSingleAgentMessageDeltaFakeAgentBackendRunClient()
|
||||
store = _FakeSessionStore()
|
||||
@ -1249,6 +1363,63 @@ def test_failed_run_raises_agent_backend_error() -> None:
|
||||
assert store.saved[0][2] == CompositorSessionSnapshot(layers=[])
|
||||
|
||||
|
||||
def test_failed_run_persists_partial_usage(sqlite_session: Session) -> None:
|
||||
sqlite_session.add(_message_record())
|
||||
sqlite_session.flush()
|
||||
|
||||
with pytest.raises(AgentBackendRunFailedError, match="failed after model calls"):
|
||||
_run(_runner(_UsageFailedClient(), _FakeSessionStore()), _FakeQueueManager())
|
||||
|
||||
sqlite_session.expire_all()
|
||||
message = sqlite_session.get(Message, "msg-1")
|
||||
assert message is not None
|
||||
assert message.message_tokens == 7
|
||||
assert message.answer_tokens == 5
|
||||
assert message.total_price == Decimal("0.12")
|
||||
assert message.provider_response_latency == 0.4
|
||||
assert message.message_metadata is not None
|
||||
assert json.loads(message.message_metadata)["usage"]["total_tokens"] == 12
|
||||
|
||||
|
||||
def test_partial_usage_persistence_ignores_missing_message() -> None:
|
||||
AgentAppRunner._persist_message_usage(
|
||||
message_id="missing",
|
||||
usage=LLMUsage.from_metadata({"prompt_tokens": 2, "completion_tokens": 1}),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("metadata", ["{", "[]"])
|
||||
def test_partial_usage_persistence_recovers_invalid_metadata(metadata: str, sqlite_session: Session) -> None:
|
||||
message = _message_record()
|
||||
message.message_metadata = metadata
|
||||
sqlite_session.add(message)
|
||||
sqlite_session.flush()
|
||||
|
||||
AgentAppRunner._persist_message_usage(
|
||||
message_id=message.id,
|
||||
usage=LLMUsage.from_metadata({"prompt_tokens": 2, "completion_tokens": 1}),
|
||||
)
|
||||
|
||||
sqlite_session.expire_all()
|
||||
persisted = sqlite_session.get(Message, message.id)
|
||||
assert persisted is not None
|
||||
assert persisted.message_metadata is not None
|
||||
assert json.loads(persisted.message_metadata)["usage"]["total_tokens"] == 3
|
||||
|
||||
|
||||
def test_partial_usage_persistence_rolls_back_database_error(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
session = MagicMock()
|
||||
session.get.side_effect = RuntimeError("database unavailable")
|
||||
monkeypatch.setattr(app_runner_module.db, "session", session)
|
||||
|
||||
AgentAppRunner._persist_message_usage(
|
||||
message_id="msg-1",
|
||||
usage=LLMUsage.from_metadata({"prompt_tokens": 2, "completion_tokens": 1}),
|
||||
)
|
||||
|
||||
session.rollback.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("outcome", ["failed", "stopped"])
|
||||
def test_snapshot_save_failure_preserves_original_app_outcome(outcome: str) -> None:
|
||||
store = _ExplodingSessionStore()
|
||||
@ -1360,6 +1531,24 @@ def test_stopped_task_waits_for_cancelled_snapshot_and_saves_session() -> None:
|
||||
assert store.saved[0][2] == CompositorSessionSnapshot(layers=[])
|
||||
|
||||
|
||||
def test_stopped_task_persists_partial_usage(sqlite_session: Session) -> None:
|
||||
sqlite_session.add(_message_record())
|
||||
sqlite_session.flush()
|
||||
client = _UsageCancellationClient()
|
||||
|
||||
with pytest.raises(GenerateTaskStoppedError):
|
||||
_run(_runner(client, _FakeSessionStore()), _StoppedQueueManager())
|
||||
|
||||
sqlite_session.expire_all()
|
||||
message = sqlite_session.get(Message, "msg-1")
|
||||
assert message is not None
|
||||
assert message.message_tokens == 13
|
||||
assert message.answer_tokens == 8
|
||||
assert message.total_price == Decimal("0.21")
|
||||
assert message.message_metadata is not None
|
||||
assert json.loads(message.message_metadata)["usage"]["total_tokens"] == 21
|
||||
|
||||
|
||||
def test_cancel_and_wait_failure_preserves_stopped_app_outcome() -> None:
|
||||
client = _CancelAndWaitFailingClient()
|
||||
store = _FakeSessionStore()
|
||||
@ -1382,7 +1571,7 @@ def test_ask_human_pauses_turn_creates_form_and_persists_correlation() -> None:
|
||||
# ENG-635/637: the PAUSED scenario emits a dify.ask_human deferred call, so
|
||||
# the chat turn ends by creating a conversation-owned HITL form + saving the
|
||||
# pause correlation, instead of crashing. Stub the form repo (DB-free).
|
||||
client = FakeAgentBackendRunClient(scenario=FakeAgentBackendScenario.PAUSED)
|
||||
client = _UsagePausedClient()
|
||||
store = _FakeSessionStore()
|
||||
qm = _FakeQueueManager()
|
||||
runner = _runner(client, store)
|
||||
@ -1400,6 +1589,7 @@ def test_ask_human_pauses_turn_creates_form_and_persists_correlation() -> None:
|
||||
assert created_params.workflow_execution_id is None
|
||||
assert [e for e in qm.events if isinstance(e, QueueMessageEndEvent)]
|
||||
assert _saved_user_query(qm) == "hello"
|
||||
assert _llm_result(qm).usage.total_tokens == 8
|
||||
# The pause correlation is persisted so a form submission can resume the run.
|
||||
assert store.saved
|
||||
assert store.saved[0][3] == "form-1"
|
||||
|
||||
@ -1,7 +1,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Generator, Sequence
|
||||
from datetime import UTC, datetime
|
||||
from decimal import Decimal
|
||||
from threading import Thread
|
||||
from typing import cast
|
||||
from unittest.mock import Mock
|
||||
@ -1276,6 +1278,82 @@ class TestEasyUiBasedGenerateTaskPipeline:
|
||||
assert trace_task.kwargs["trace_session_id"] == "session-1"
|
||||
assert len(sent_payloads) == 1
|
||||
|
||||
def test_save_stopped_message_preserves_backend_reported_usage(
|
||||
self, monkeypatch: pytest.MonkeyPatch, sqlite_session: Session
|
||||
) -> None:
|
||||
pipeline = EasyUIBasedGenerateTaskPipeline(
|
||||
application_generate_entity=_make_entity(ChatAppGenerateEntity, AppMode.CHAT),
|
||||
queue_manager=_FakeQueueManager(),
|
||||
conversation=_make_conversation(AppMode.CHAT),
|
||||
message=_make_message(),
|
||||
stream=True,
|
||||
)
|
||||
_set_method(pipeline, "_model_config", _ModelConfigMode(mode="chat"))
|
||||
pipeline._task_state.llm_result.usage = LLMUsage.empty_usage()
|
||||
|
||||
message = _make_message()
|
||||
message.message_tokens = 13
|
||||
message.answer_tokens = 8
|
||||
message.total_price = Decimal("0.21")
|
||||
message.message_metadata = json.dumps({"usage": {"total_tokens": 21}})
|
||||
sqlite_session.add_all([_make_conversation(AppMode.CHAT), message])
|
||||
sqlite_session.flush()
|
||||
monkeypatch.setattr(
|
||||
"core.app.task_pipeline.easy_ui_based_generate_task_pipeline.PromptMessageUtil.prompt_messages_to_prompt_for_saving",
|
||||
lambda _mode, _prompt_messages: "",
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"core.app.task_pipeline.easy_ui_based_generate_task_pipeline.message_was_created.send",
|
||||
lambda *_args, **_kwargs: None,
|
||||
)
|
||||
|
||||
pipeline._save_message(session=sqlite_session, preserve_existing_usage=True)
|
||||
|
||||
assert message.message_tokens == 13
|
||||
assert message.answer_tokens == 8
|
||||
assert message.total_price == Decimal("0.21")
|
||||
assert message.message_metadata is not None
|
||||
assert json.loads(message.message_metadata)["usage"]["total_tokens"] == 21
|
||||
|
||||
@pytest.mark.parametrize("metadata", ["{", "[]"])
|
||||
def test_save_stopped_message_recovers_invalid_existing_metadata(
|
||||
self,
|
||||
metadata: str,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
pipeline = EasyUIBasedGenerateTaskPipeline(
|
||||
application_generate_entity=_make_entity(ChatAppGenerateEntity, AppMode.CHAT),
|
||||
queue_manager=_FakeQueueManager(),
|
||||
conversation=_make_conversation(AppMode.CHAT),
|
||||
message=_make_message(),
|
||||
stream=True,
|
||||
)
|
||||
_set_method(pipeline, "_model_config", _ModelConfigMode(mode="chat"))
|
||||
pipeline._task_state.llm_result.usage = LLMUsage.empty_usage()
|
||||
|
||||
message = _make_message()
|
||||
message.message_tokens = 13
|
||||
message.answer_tokens = 8
|
||||
message.total_price = Decimal("0.21")
|
||||
message.message_metadata = metadata
|
||||
sqlite_session.add_all([_make_conversation(AppMode.CHAT), message])
|
||||
sqlite_session.flush()
|
||||
monkeypatch.setattr(
|
||||
"core.app.task_pipeline.easy_ui_based_generate_task_pipeline.PromptMessageUtil.prompt_messages_to_prompt_for_saving",
|
||||
lambda _mode, _prompt_messages: "",
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"core.app.task_pipeline.easy_ui_based_generate_task_pipeline.message_was_created.send",
|
||||
lambda *_args, **_kwargs: None,
|
||||
)
|
||||
|
||||
pipeline._save_message(session=sqlite_session, preserve_existing_usage=True)
|
||||
|
||||
assert message.message_tokens == 13
|
||||
assert message.answer_tokens == 8
|
||||
assert message.total_price == Decimal("0.21")
|
||||
|
||||
def test_save_message_raises_when_message_not_found(self, sqlite_session: Session):
|
||||
conversation = _make_conversation(AppMode.CHAT)
|
||||
message = _make_message()
|
||||
|
||||
@ -152,6 +152,14 @@ def test_failure_output_adapter_preserves_backend_failed_reason():
|
||||
source_event_id="2-0",
|
||||
error="bad request",
|
||||
reason="validation",
|
||||
session_snapshot=CompositorSessionSnapshot(layers=[]),
|
||||
usage={
|
||||
"prompt_tokens": 13,
|
||||
"completion_tokens": 8,
|
||||
"total_tokens": 21,
|
||||
"total_price": "0.000210",
|
||||
"currency": "USD",
|
||||
},
|
||||
),
|
||||
inputs={},
|
||||
process_data={},
|
||||
@ -161,6 +169,12 @@ def test_failure_output_adapter_preserves_backend_failed_reason():
|
||||
assert result.status == WorkflowNodeExecutionStatus.FAILED
|
||||
assert result.error == "bad request"
|
||||
assert result.error_type == "validation"
|
||||
assert result.llm_usage.total_tokens == 21
|
||||
assert result.metadata[WorkflowNodeExecutionMetadataKey.TOTAL_TOKENS] == 21
|
||||
assert result.metadata[WorkflowNodeExecutionMetadataKey.TOTAL_PRICE] == Decimal("0.000210")
|
||||
assert result.metadata[WorkflowNodeExecutionMetadataKey.AGENT_LOG]["agent_backend"]["session_snapshot"] == {
|
||||
"layer_count": 0
|
||||
}
|
||||
|
||||
|
||||
def test_failure_output_adapter_prefers_run_failure_type_over_reason():
|
||||
@ -596,6 +610,7 @@ def test_failure_output_adapter_maps_cancelled_to_failure_code():
|
||||
source_event_id="2-0",
|
||||
reason="user_cancelled",
|
||||
message=None,
|
||||
usage={"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8},
|
||||
),
|
||||
inputs={},
|
||||
process_data={},
|
||||
@ -604,6 +619,8 @@ def test_failure_output_adapter_maps_cancelled_to_failure_code():
|
||||
|
||||
assert result.status == WorkflowNodeExecutionStatus.FAILED
|
||||
assert result.error_type == "agent_backend_run_cancelled"
|
||||
assert result.llm_usage.total_tokens == 8
|
||||
assert result.metadata[WorkflowNodeExecutionMetadataKey.TOTAL_TOKENS] == 8
|
||||
|
||||
|
||||
def test_stream_exhausted_result_is_failed_with_stream_error():
|
||||
|
||||
@ -331,6 +331,7 @@ class RunFailedEventData(BaseModel):
|
||||
error_type: RunFailureType | None = None
|
||||
reason: str | None = None
|
||||
session_snapshot: CompositorSessionSnapshot | None = None
|
||||
usage: AgentRunUsage | None = None
|
||||
|
||||
model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid")
|
||||
|
||||
@ -341,6 +342,7 @@ class RunCancelledEventData(BaseModel):
|
||||
reason: str | None = None
|
||||
message: str | None = None
|
||||
session_snapshot: CompositorSessionSnapshot | None = None
|
||||
usage: AgentRunUsage | None = None
|
||||
|
||||
model_config: ClassVar[ConfigDict] = ConfigDict(extra="forbid")
|
||||
|
||||
|
||||
@ -186,6 +186,7 @@ async def emit_run_failed(
|
||||
error_type: RunFailureType | None = None,
|
||||
reason: str | None = None,
|
||||
session_snapshot: CompositorSessionSnapshot | None = None,
|
||||
usage: AgentRunUsage | None = None,
|
||||
) -> RunFinalizationResult:
|
||||
"""Finalize a run with a failed terminal event."""
|
||||
return await sink.finalize_run(
|
||||
@ -196,6 +197,7 @@ async def emit_run_failed(
|
||||
error_type=error_type,
|
||||
reason=reason,
|
||||
session_snapshot=session_snapshot,
|
||||
usage=usage,
|
||||
),
|
||||
created_at=utc_now(),
|
||||
),
|
||||
|
||||
@ -21,7 +21,7 @@ from typing import Protocol
|
||||
import httpx
|
||||
|
||||
from agenton.compositor import CompositorSessionSnapshot, LayerProviderInput
|
||||
from dify_agent.protocol.schemas import CancelRunRequest, CancelRunResponse, CreateRunRequest, RunStatus
|
||||
from dify_agent.protocol.schemas import AgentRunUsage, CancelRunRequest, CancelRunResponse, CreateRunRequest, RunStatus
|
||||
from dify_agent.runtime.cancellation import RunCancellationIntent
|
||||
from dify_agent.runtime.compositor_factory import create_default_layer_providers
|
||||
from dify_agent.runtime.event_sink import RunEventSink, RunFinalizationResult, emit_run_failed
|
||||
@ -64,6 +64,7 @@ class RunStore(RunEventSink, Protocol):
|
||||
intent: RunCancellationIntent,
|
||||
*,
|
||||
session_snapshot: CompositorSessionSnapshot | None = None,
|
||||
usage: AgentRunUsage | None = None,
|
||||
) -> RunFinalizationResult:
|
||||
"""Publish cancellation after the owner runner has exited."""
|
||||
...
|
||||
@ -77,6 +78,11 @@ class RunnableRun(Protocol):
|
||||
"""Return the post-exit snapshot for the current invocation, if available."""
|
||||
...
|
||||
|
||||
@property
|
||||
def terminal_usage(self) -> AgentRunUsage | None:
|
||||
"""Return usage accumulated before this run exited, if available."""
|
||||
...
|
||||
|
||||
async def run(self) -> None:
|
||||
"""Run until terminal status/events have been written or cancellation occurs."""
|
||||
...
|
||||
@ -198,6 +204,7 @@ class RunScheduler:
|
||||
error=f"run cancellation observer failed: {exc}",
|
||||
reason="cancellation_observer",
|
||||
session_snapshot=runner.terminal_session_snapshot,
|
||||
usage=runner.terminal_usage,
|
||||
)
|
||||
if not finalization.applied and finalization.status == "running":
|
||||
intent = await self.store.get_cancellation_intent(record.run_id)
|
||||
@ -206,6 +213,7 @@ class RunScheduler:
|
||||
record.run_id,
|
||||
intent,
|
||||
session_snapshot=runner.terminal_session_snapshot,
|
||||
usage=runner.terminal_usage,
|
||||
)
|
||||
raise
|
||||
|
||||
@ -214,6 +222,7 @@ class RunScheduler:
|
||||
record.run_id,
|
||||
intent,
|
||||
session_snapshot=runner.terminal_session_snapshot,
|
||||
usage=runner.terminal_usage,
|
||||
)
|
||||
else:
|
||||
runner_error: Exception | None = None
|
||||
@ -228,6 +237,7 @@ class RunScheduler:
|
||||
record.run_id,
|
||||
intent,
|
||||
session_snapshot=runner.terminal_session_snapshot,
|
||||
usage=runner.terminal_usage,
|
||||
)
|
||||
if runner_error is not None:
|
||||
raise runner_error
|
||||
@ -240,11 +250,13 @@ class RunScheduler:
|
||||
record.run_id,
|
||||
intent,
|
||||
session_snapshot=runner.terminal_session_snapshot,
|
||||
usage=runner.terminal_usage,
|
||||
)
|
||||
else:
|
||||
finalization = await self._mark_cancelled_run_failed(
|
||||
record.run_id,
|
||||
session_snapshot=runner.terminal_session_snapshot,
|
||||
usage=runner.terminal_usage,
|
||||
)
|
||||
if finalization is not None and not finalization.applied and finalization.status == "running":
|
||||
intent = await self.store.get_cancellation_intent(record.run_id)
|
||||
@ -253,6 +265,7 @@ class RunScheduler:
|
||||
record.run_id,
|
||||
intent,
|
||||
session_snapshot=runner.terminal_session_snapshot,
|
||||
usage=runner.terminal_usage,
|
||||
)
|
||||
raise
|
||||
except Exception:
|
||||
@ -308,6 +321,7 @@ class RunScheduler:
|
||||
run_id: str,
|
||||
*,
|
||||
session_snapshot: CompositorSessionSnapshot | None = None,
|
||||
usage: AgentRunUsage | None = None,
|
||||
) -> RunFinalizationResult | None:
|
||||
"""Best-effort failure event/status for shutdown-cancelled runs."""
|
||||
message = "run cancelled during server shutdown"
|
||||
@ -318,6 +332,7 @@ class RunScheduler:
|
||||
error=message,
|
||||
reason="shutdown",
|
||||
session_snapshot=session_snapshot,
|
||||
usage=usage,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("failed to mark cancelled run failed", extra={"run_id": run_id})
|
||||
|
||||
@ -185,6 +185,7 @@ class AgentRunRunner:
|
||||
is_cancelled: Callable[[], bool]
|
||||
run_timeout_seconds: float
|
||||
_terminal_session_snapshot: CompositorSessionSnapshot | None
|
||||
_terminal_usage: AgentRunUsage | None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@ -207,15 +208,22 @@ class AgentRunRunner:
|
||||
self.is_cancelled = is_cancelled or (lambda: False)
|
||||
self.run_timeout_seconds = run_timeout_seconds
|
||||
self._terminal_session_snapshot = None
|
||||
self._terminal_usage = None
|
||||
|
||||
@property
|
||||
def terminal_session_snapshot(self) -> CompositorSessionSnapshot | None:
|
||||
"""Return the snapshot captured after the current compositor context exited."""
|
||||
return self._terminal_session_snapshot
|
||||
|
||||
@property
|
||||
def terminal_usage(self) -> AgentRunUsage | None:
|
||||
"""Return usage accumulated before the current run reached any terminal state."""
|
||||
return self._terminal_usage
|
||||
|
||||
async def run(self) -> None:
|
||||
"""Execute the run and emit the documented event sequence."""
|
||||
self._terminal_session_snapshot = None
|
||||
self._terminal_usage = None
|
||||
if self.is_cancelled():
|
||||
return
|
||||
_ = await emit_run_started(self.sink, run_id=self.run_id)
|
||||
@ -233,6 +241,7 @@ class AgentRunRunner:
|
||||
error_type=error_type,
|
||||
reason=reason,
|
||||
session_snapshot=self._terminal_session_snapshot,
|
||||
usage=self._terminal_usage,
|
||||
)
|
||||
if finalization.applied:
|
||||
raise
|
||||
@ -296,6 +305,7 @@ class AgentRunRunner:
|
||||
deferred_tool_call: DeferredToolCallPayload | None = None
|
||||
result_kind: Literal["output", "deferred_tool_call"] | None = None
|
||||
usage: AgentRunUsage | None = None
|
||||
model: Any = None
|
||||
run = None
|
||||
try:
|
||||
async with compositor.enter(configs=layer_configs, session_snapshot=self.request.session_snapshot) as run:
|
||||
@ -370,6 +380,7 @@ class AgentRunRunner:
|
||||
) from exc
|
||||
complete_usage = model.accumulated_usage if isinstance(model, _HasAccumulatedUsage) else None
|
||||
usage = _serialize_agent_usage(complete_usage if complete_usage is not None else _result_usage(result))
|
||||
self._terminal_usage = usage
|
||||
replace_successful_run_history(history_layer, result.all_messages())
|
||||
if isinstance(result.output, DeferredToolRequests):
|
||||
if ask_human_layer is None:
|
||||
@ -396,6 +407,10 @@ class AgentRunRunner:
|
||||
finally:
|
||||
if entered_run and run is not None:
|
||||
self._terminal_session_snapshot = run.session_snapshot
|
||||
if isinstance(model, _HasAccumulatedUsage):
|
||||
accumulated_usage = _serialize_agent_usage(model.accumulated_usage)
|
||||
if accumulated_usage is not None:
|
||||
self._terminal_usage = accumulated_usage
|
||||
|
||||
if run is None or run.session_snapshot is None:
|
||||
raise RuntimeError("Agenton run did not produce a session snapshot after exit.")
|
||||
|
||||
@ -16,6 +16,7 @@ from redis.asyncio import Redis
|
||||
|
||||
from agenton.compositor import CompositorSessionSnapshot
|
||||
from dify_agent.protocol.schemas import (
|
||||
AgentRunUsage,
|
||||
RUN_EVENT_ADAPTER,
|
||||
CancelRunRequest,
|
||||
RunCancelledEvent,
|
||||
@ -287,6 +288,7 @@ class RedisRunStore(RunEventSink):
|
||||
intent: RunCancellationIntent,
|
||||
*,
|
||||
session_snapshot: CompositorSessionSnapshot | None = None,
|
||||
usage: AgentRunUsage | None = None,
|
||||
) -> RunFinalizationResult:
|
||||
"""Atomically publish cancellation after the owner runner has exited."""
|
||||
event = RunCancelledEvent(
|
||||
@ -295,6 +297,7 @@ class RedisRunStore(RunEventSink):
|
||||
reason=intent.reason,
|
||||
message=intent.message,
|
||||
session_snapshot=session_snapshot,
|
||||
usage=usage,
|
||||
),
|
||||
created_at=utc_now(),
|
||||
)
|
||||
|
||||
@ -489,6 +489,37 @@ def test_run_succeeded_event_round_trips_usage() -> None:
|
||||
assert b'"usage"' in payload
|
||||
|
||||
|
||||
def test_run_failed_event_round_trips_usage() -> None:
|
||||
usage = AgentRunUsage(prompt_tokens=13, completion_tokens=8)
|
||||
event = RunFailedEvent(run_id="run-partial-usage", data=RunFailedEventData(error="boom", usage=usage))
|
||||
|
||||
payload = RUN_EVENT_ADAPTER.dump_json(event)
|
||||
decoded = RUN_EVENT_ADAPTER.validate_json(payload)
|
||||
|
||||
assert isinstance(decoded, RunFailedEvent)
|
||||
assert decoded.data.usage is not None
|
||||
assert decoded.data.usage.prompt_tokens == 13
|
||||
assert decoded.data.usage.completion_tokens == 8
|
||||
assert decoded.data.usage.total_tokens == 21
|
||||
|
||||
|
||||
def test_run_cancelled_event_round_trips_usage() -> None:
|
||||
usage = AgentRunUsage(prompt_tokens=13, completion_tokens=8)
|
||||
event = RunCancelledEvent(
|
||||
run_id="run-partial-usage",
|
||||
data=RunCancelledEventData(reason="user_cancelled", usage=usage),
|
||||
)
|
||||
|
||||
payload = RUN_EVENT_ADAPTER.dump_json(event)
|
||||
decoded = RUN_EVENT_ADAPTER.validate_json(payload)
|
||||
|
||||
assert isinstance(decoded, RunCancelledEvent)
|
||||
assert decoded.data.usage is not None
|
||||
assert decoded.data.usage.prompt_tokens == 13
|
||||
assert decoded.data.usage.completion_tokens == 8
|
||||
assert decoded.data.usage.total_tokens == 21
|
||||
|
||||
|
||||
def test_run_succeeded_event_round_trips_complete_pricing_usage() -> None:
|
||||
event = RunSucceededEvent(
|
||||
run_id="run-priced-usage",
|
||||
|
||||
@ -13,6 +13,7 @@ from dify_agent.layers.execution_context import DIFY_EXECUTION_CONTEXT_LAYER_TYP
|
||||
from dify_agent.layers.output import DIFY_OUTPUT_LAYER_TYPE_ID, DifyOutputLayerConfig
|
||||
from dify_agent.protocol import DIFY_AGENT_MODEL_LAYER_ID, DIFY_AGENT_OUTPUT_LAYER_ID, RunFailureType
|
||||
from dify_agent.protocol.schemas import (
|
||||
AgentRunUsage,
|
||||
CancelRunRequest,
|
||||
CreateRunRequest,
|
||||
RunCancelledEvent,
|
||||
@ -174,6 +175,7 @@ class FakeStore:
|
||||
intent: RunCancellationIntent,
|
||||
*,
|
||||
session_snapshot: CompositorSessionSnapshot | None = None,
|
||||
usage: AgentRunUsage | None = None,
|
||||
) -> RunFinalizationResult:
|
||||
current_status = self.statuses[run_id]
|
||||
if current_status != "running":
|
||||
@ -186,6 +188,7 @@ class FakeStore:
|
||||
reason=intent.reason,
|
||||
message=intent.message,
|
||||
session_snapshot=session_snapshot,
|
||||
usage=usage,
|
||||
),
|
||||
)
|
||||
event_id = str(len(self.events[run_id]) + 1)
|
||||
@ -267,12 +270,17 @@ class SnapshotlessRunner:
|
||||
def terminal_session_snapshot(self) -> CompositorSessionSnapshot | None:
|
||||
return None
|
||||
|
||||
@property
|
||||
def terminal_usage(self) -> AgentRunUsage | None:
|
||||
return None
|
||||
|
||||
|
||||
class ControlledRunner:
|
||||
started: asyncio.Event
|
||||
release: asyncio.Event
|
||||
finished: asyncio.Event | None
|
||||
_terminal_session_snapshot: CompositorSessionSnapshot
|
||||
_terminal_usage: AgentRunUsage | None
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@ -280,16 +288,22 @@ class ControlledRunner:
|
||||
started: asyncio.Event,
|
||||
release: asyncio.Event,
|
||||
finished: asyncio.Event | None = None,
|
||||
usage: AgentRunUsage | None = None,
|
||||
) -> None:
|
||||
self.started = started
|
||||
self.release = release
|
||||
self.finished = finished
|
||||
self._terminal_session_snapshot = CompositorSessionSnapshot(layers=[])
|
||||
self._terminal_usage = usage
|
||||
|
||||
@property
|
||||
def terminal_session_snapshot(self) -> CompositorSessionSnapshot:
|
||||
return self._terminal_session_snapshot
|
||||
|
||||
@property
|
||||
def terminal_usage(self) -> AgentRunUsage | None:
|
||||
return self._terminal_usage
|
||||
|
||||
async def run(self) -> None:
|
||||
_ = self.started.set()
|
||||
try:
|
||||
@ -639,6 +653,7 @@ def test_non_owner_cancel_run_stops_owner_task_and_persists_cancelled_terminal()
|
||||
started=started,
|
||||
release=asyncio.Event(),
|
||||
finished=runner_finished,
|
||||
usage=AgentRunUsage(prompt_tokens=13, completion_tokens=8),
|
||||
),
|
||||
)
|
||||
remote_scheduler = RunScheduler(
|
||||
@ -665,6 +680,10 @@ def test_non_owner_cancel_run_stops_owner_task_and_persists_cancelled_terminal()
|
||||
terminal = store.events[record.run_id][0]
|
||||
assert isinstance(terminal, RunCancelledEvent)
|
||||
assert terminal.data.session_snapshot == CompositorSessionSnapshot(layers=[])
|
||||
assert terminal.data.usage is not None
|
||||
assert terminal.data.usage.prompt_tokens == 13
|
||||
assert terminal.data.usage.completion_tokens == 8
|
||||
assert terminal.data.usage.total_tokens == 21
|
||||
assert runner_finished.is_set()
|
||||
assert store.observer_finished.is_set()
|
||||
await asyncio.sleep(0)
|
||||
|
||||
@ -837,10 +837,23 @@ def test_runner_timeout_cancels_agent_and_releases_runtime_lease(monkeypatch: py
|
||||
agent_cancelled = False
|
||||
shell_client = FakeRunnerShellctlClient()
|
||||
|
||||
class FakeUsageModel(TestModel):
|
||||
@property
|
||||
def accumulated_usage(self) -> LLMUsage:
|
||||
return LLMUsage.from_metadata(
|
||||
{
|
||||
"prompt_tokens": 13,
|
||||
"completion_tokens": 8,
|
||||
"total_tokens": 21,
|
||||
"total_price": "0.000210",
|
||||
"currency": "USD",
|
||||
}
|
||||
)
|
||||
|
||||
def fake_get_model(_self: DifyPluginLLMLayer, *, http_client: httpx.AsyncClient, agent_run_id: str):
|
||||
assert http_client.is_closed is False
|
||||
assert agent_run_id == "run-timeout"
|
||||
return TestModel(custom_output_text="unused") # pyright: ignore[reportReturnType]
|
||||
return FakeUsageModel(custom_output_text="unused") # pyright: ignore[reportReturnType]
|
||||
|
||||
class FakeAgent:
|
||||
async def run(self, *_args: object, **_kwargs: object) -> None:
|
||||
@ -877,6 +890,11 @@ def test_runner_timeout_cancels_agent_and_releases_runtime_lease(monkeypatch: py
|
||||
terminal = sink.events["run-timeout"][-1]
|
||||
assert isinstance(terminal, RunFailedEvent)
|
||||
assert terminal.data.error_type is RunFailureType.AGENT_RUN_LIMIT_EXCEEDED
|
||||
assert terminal.data.usage is not None
|
||||
assert terminal.data.usage.prompt_tokens == 13
|
||||
assert terminal.data.usage.completion_tokens == 8
|
||||
assert terminal.data.usage.total_tokens == 21
|
||||
assert terminal.data.usage.total_price == Decimal("0.000210")
|
||||
assert sink.statuses["run-timeout"] == "failed"
|
||||
assert agent_cancelled is True
|
||||
assert shell_client.closed is True
|
||||
|
||||
@ -9,6 +9,7 @@ from pydantic import JsonValue
|
||||
from agenton.compositor import CompositorSessionSnapshot, LayerSessionSnapshot
|
||||
from agenton.layers import LifecycleState
|
||||
from dify_agent.protocol.schemas import (
|
||||
AgentRunUsage,
|
||||
RUN_EVENT_ADAPTER,
|
||||
CancelRunRequest,
|
||||
RunCancelledEvent,
|
||||
@ -239,6 +240,7 @@ def test_finalize_cancellation_maps_eval_result_and_arguments() -> None:
|
||||
"run-1",
|
||||
intent,
|
||||
session_snapshot=CompositorSessionSnapshot(layers=[]),
|
||||
usage=AgentRunUsage(prompt_tokens=13, completion_tokens=8),
|
||||
)
|
||||
)
|
||||
|
||||
@ -253,11 +255,12 @@ def test_finalize_cancellation_maps_eval_result_and_arguments() -> None:
|
||||
payload = json.loads(cast(str, eval_command[9]))
|
||||
assert "id" not in payload
|
||||
assert payload["type"] == "run_cancelled"
|
||||
assert payload["data"] == {
|
||||
"reason": "workflow_aborted",
|
||||
"message": "workflow stopped",
|
||||
"session_snapshot": {"schema_version": 1, "layers": []},
|
||||
}
|
||||
assert payload["data"]["reason"] == "workflow_aborted"
|
||||
assert payload["data"]["message"] == "workflow stopped"
|
||||
assert payload["data"]["session_snapshot"] == {"schema_version": 1, "layers": []}
|
||||
assert payload["data"]["usage"]["prompt_tokens"] == 13
|
||||
assert payload["data"]["usage"]["completion_tokens"] == 8
|
||||
assert payload["data"]["usage"]["total_tokens"] == 21
|
||||
assert eval_command[10] == "60"
|
||||
|
||||
|
||||
|
||||
Loading…
Reference in New Issue
Block a user