diff --git a/api/core/app/llm/quota.py b/api/core/app/llm/quota.py index e4e84502cfa..762c435ca58 100644 --- a/api/core/app/llm/quota.py +++ b/api/core/app/llm/quota.py @@ -111,7 +111,13 @@ def _get_current_quota_configuration(system_configuration): ) -def reserve_llm_quota_for_model(*, tenant_id: str, provider: str, model: str) -> LLMQuotaReservation: +def reserve_llm_quota_for_model( + *, + tenant_id: str, + provider: str, + model: str, + request_id: str | None = None, +) -> LLMQuotaReservation: """Reserve system-hosted LLM quota before invoking the provider.""" provider_configuration = _get_provider_configuration(tenant_id=tenant_id, provider=provider) reservation = LLMQuotaReservation( @@ -151,7 +157,7 @@ def reserve_llm_quota_for_model(*, tenant_id: str, provider: str, model: str) -> tenant_id=tenant_id, credits_required=amount, pool_type="paid" if quota_type == ProviderQuotaType.PAID else "trial", - request_id=str(uuid4()), + request_id=request_id or str(uuid4()), session_factory=db.session, meta={"source": "llm.invoke", "provider": provider, "model": model}, ) diff --git a/api/core/model_manager.py b/api/core/model_manager.py index c07cc74583d..5a9cab1afe2 100644 --- a/api/core/model_manager.py +++ b/api/core/model_manager.py @@ -2,6 +2,7 @@ import logging from collections.abc import Callable, Generator, Iterable, Mapping, Sequence from copy import deepcopy from typing import IO, Any, Literal, Optional, ParamSpec, TypeVar, Union, cast, overload, override +from uuid import UUID from configs import dify_config from core.entities import PluginCredentialType @@ -445,15 +446,32 @@ class ModelInstance: class QuotaManagedModelInstance(ModelInstance): """A system-hosted LLM instance that owns quota settlement per invocation.""" - def reserve_quota(self): + def reserve_quota(self, *, request_id: str | None = None): from core.app.llm.quota import reserve_llm_quota_for_model return reserve_llm_quota_for_model( tenant_id=self.provider_model_bundle.configuration.tenant_id, provider=self.provider, model=self.model_name, + request_id=request_id, ) + @staticmethod + def _get_reservation_request_id(request_metadata: Mapping[str, object] | None) -> str | None: + request_id = request_metadata.get("invocation_id") if request_metadata else None + if not isinstance(request_id, str) or not request_id: + return None + try: + return str(UUID(request_id)) + except ValueError: + return None + + def _reserve_quota_for_request(self, request_metadata: Mapping[str, object] | None): + request_id = self._get_reservation_request_id(request_metadata) + if request_id is None: + return self.reserve_quota() + return self.reserve_quota(request_id=request_id) + @staticmethod def release_quota_safely(reservation) -> None: try: @@ -520,7 +538,7 @@ class QuotaManagedModelInstance(ModelInstance): request_metadata=request_metadata, ) - reservation = self.reserve_quota() + reservation = self._reserve_quota_for_request(request_metadata) try: response = super().invoke_llm( prompt_messages=normalized_prompt_messages, @@ -548,7 +566,7 @@ class QuotaManagedModelInstance(ModelInstance): callbacks: list[Callback] | None, request_metadata: Mapping[str, object] | None, ) -> Generator: - reservation = self.reserve_quota() + reservation = self._reserve_quota_for_request(request_metadata) usage: LLMUsage | None = None try: response = super().invoke_llm( diff --git a/api/tests/unit_tests/core/app/test_llm_quota.py b/api/tests/unit_tests/core/app/test_llm_quota.py index e7297639610..631bde3d989 100644 --- a/api/tests/unit_tests/core/app/test_llm_quota.py +++ b/api/tests/unit_tests/core/app/test_llm_quota.py @@ -2,6 +2,7 @@ from collections.abc import Generator from contextlib import contextmanager from types import SimpleNamespace from unittest.mock import ANY, MagicMock, patch +from uuid import UUID import pytest from sqlalchemy import create_engine, select @@ -130,6 +131,7 @@ def test_reserve_llm_quota_uses_exact_credit_pool_reservation() -> None: tenant_id="tenant-id", provider="openai", model="gpt-4o", + request_id="11111111-1111-5111-8111-111111111111", ) reservation.commit(LLMUsage.empty_usage()) reservation.release() @@ -140,7 +142,7 @@ def test_reserve_llm_quota_uses_exact_credit_pool_reservation() -> None: tenant_id="tenant-id", credits_required=9, pool_type="trial", - request_id=ANY, + request_id="11111111-1111-5111-8111-111111111111", session_factory=ANY, meta={"source": "llm.invoke", "provider": "openai", "model": "gpt-4o"}, ) @@ -148,6 +150,35 @@ def test_reserve_llm_quota_uses_exact_credit_pool_reservation() -> None: credit_reservation.release.assert_not_called() +def test_reserve_llm_quota_generates_request_id_when_not_supplied() -> None: + credit_reservation = MagicMock() + provider_configuration = SimpleNamespace( + using_provider_type=ProviderType.SYSTEM, + get_provider_model=MagicMock(return_value=SimpleNamespace(status=ModelStatus.ACTIVE)), + system_configuration=SimpleNamespace( + current_quota_type=ProviderQuotaType.TRIAL, + quota_configurations=[ + SimpleNamespace( + quota_type=ProviderQuotaType.TRIAL, + quota_unit=QuotaUnit.TIMES, + quota_limit=100, + ) + ], + ), + ) + provider_manager = MagicMock() + provider_manager.get_configurations.return_value.get.return_value = provider_configuration + + with ( + patch("core.app.llm.quota.create_plugin_provider_manager", return_value=provider_manager), + patch("core.app.llm.quota.CreditPoolService.reserve_credits", return_value=credit_reservation) as reserve, + ): + reserve_llm_quota_for_model(tenant_id="tenant-id", provider="openai", model="gpt-4o") + + generated_request_id = reserve.call_args.kwargs["request_id"] + assert str(UUID(generated_request_id)) == generated_request_id + + def test_reserve_llm_quota_requires_accurate_usage_for_free_tokens() -> None: provider_configuration = SimpleNamespace( using_provider_type=ProviderType.SYSTEM, diff --git a/api/tests/unit_tests/core/test_model_manager.py b/api/tests/unit_tests/core/test_model_manager.py index fa873e8d865..3874321093c 100644 --- a/api/tests/unit_tests/core/test_model_manager.py +++ b/api/tests/unit_tests/core/test_model_manager.py @@ -1,4 +1,5 @@ from unittest.mock import MagicMock, patch +from uuid import uuid4 import pytest import redis @@ -145,13 +146,19 @@ def test_quota_managed_non_streaming_invocation_finalizes_reservation() -> None: result = MagicMock(spec=LLMResult, usage=usage) reservation = MagicMock(commit_before_delivery=True) + invocation_id = str(uuid4()) with ( - patch.object(model_instance, "reserve_quota", return_value=reservation), + patch.object(model_instance, "reserve_quota", return_value=reservation) as reserve_quota, patch.object(ModelInstance, "invoke_llm", return_value=result) as invoke, ): - response = model_instance.invoke_llm(prompt_messages=[], stream=False) + response = model_instance.invoke_llm( + prompt_messages=[], + stream=False, + request_metadata={"invocation_id": invocation_id}, + ) assert response is result + reserve_quota.assert_called_once_with(request_id=invocation_id) invoke.assert_called_once() reservation.commit.assert_called_once_with(usage) reservation.release.assert_called_once_with() @@ -172,17 +179,23 @@ def test_quota_managed_stream_commits_before_first_chunk() -> None: events: list[str] = [] reservation.commit.side_effect = lambda _usage: events.append("commit") + invocation_id = str(uuid4()) with ( - patch.object(model_instance, "reserve_quota", return_value=reservation), + patch.object(model_instance, "reserve_quota", return_value=reservation) as reserve_quota, patch.object(ModelInstance, "invoke_llm", return_value=(item for item in [chunk])), ): - response = model_instance.invoke_llm(prompt_messages=[], stream=True) + response = model_instance.invoke_llm( + prompt_messages=[], + stream=True, + request_metadata={"invocation_id": invocation_id}, + ) assert next(response) is chunk events.append("delivered") with pytest.raises(StopIteration): next(response) assert events == ["commit", "delivered"] + reserve_quota.assert_called_once_with(request_id=invocation_id) reservation.release.assert_called_once_with() diff --git a/api/tests/unit_tests/services/test_agent_llm_inner_service.py b/api/tests/unit_tests/services/test_agent_llm_inner_service.py index d69dbcf414b..ec53ded0ba3 100644 --- a/api/tests/unit_tests/services/test_agent_llm_inner_service.py +++ b/api/tests/unit_tests/services/test_agent_llm_inner_service.py @@ -2,17 +2,21 @@ from collections.abc import Generator from decimal import Decimal +from types import SimpleNamespace from unittest.mock import MagicMock, patch from uuid import uuid4 import pytest from sqlalchemy.orm import Session, sessionmaker +from configs import dify_config from core.entities.model_entities import ModelStatus +from core.entities.provider_entities import ProviderQuotaType, QuotaUnit from core.model_manager import ModelInstance, QuotaManagedModelInstance from graphon.model_runtime.entities.llm_entities import LLMResultChunk, LLMResultChunkDelta, LLMUsage from graphon.model_runtime.entities.message_entities import AssistantPromptMessage, UserPromptMessage from models.model import App, AppMode +from models.provider import ProviderType from services.agent_llm_inner_service import AgentLLMInnerService, AgentLLMInnerServiceError, PreparedAgentLLMInvocation from services.entities.agent_llm_inner import AgentLLMInvokeCaller, AgentLLMInvokeRequest, AgentLLMInvokeTarget @@ -185,12 +189,75 @@ def test_gateway_uses_quota_managed_instance_as_single_credit_owner( chunks = list(service.invoke(prepared)) assert chunks == [provider_chunk] - model_instance.reserve_quota.assert_called_once_with() + model_instance.reserve_quota.assert_called_once_with(request_id=request.caller.invocation_id) reservation.commit.assert_called_once_with(provider_chunk.delta.usage) reservation.release.assert_called_once_with() provider_invoke.assert_called_once() +def test_retried_gateway_delivery_uses_one_effective_billing_charge( + sqlite_session_factory: sessionmaker[Session], + sqlite_session: Session, +) -> None: + request = _request() + _persist_app(sqlite_session, request=request) + service = AgentLLMInnerService(session_factory=sqlite_session_factory) + model_instance, _ = _model_instance() + model_instance.provider_model_bundle.configuration.tenant_id = request.caller.tenant_id + prepared = _prepare(service, request, model_instance) + provider_configuration = SimpleNamespace( + using_provider_type=ProviderType.SYSTEM, + get_provider_model=MagicMock(return_value=SimpleNamespace(status=ModelStatus.ACTIVE)), + system_configuration=SimpleNamespace( + current_quota_type=ProviderQuotaType.TRIAL, + quota_configurations=[ + SimpleNamespace( + quota_type=ProviderQuotaType.TRIAL, + quota_unit=QuotaUnit.CREDITS, + quota_limit=100, + ) + ], + ), + ) + reservations_by_request: dict[str, str] = {} + committed_reservations: set[str] = set() + effective_charges = 0 + reserve_request_ids: list[str] = [] + + def reserve(*, request_id: str, **_: object) -> dict[str, object]: + reserve_request_ids.append(request_id) + reservation_id = reservations_by_request.setdefault(request_id, "reservation-1") + return {"reservation_id": reservation_id, "available": 97, "reserved": 3} + + def commit(*, reservation_id: str, actual_amount: int, **_: object) -> dict[str, object]: + nonlocal effective_charges + if reservation_id not in committed_reservations: + committed_reservations.add(reservation_id) + effective_charges += actual_amount + return {"available": 97, "reserved": 0, "refunded": 0} + + def provider_invoke(*_: object, **__: object) -> Generator[LLMResultChunk, None, None]: + yield _chunk("done", usage=_usage()) + + with ( + patch("core.app.llm.quota._get_provider_configuration", return_value=provider_configuration), + patch.object(type(dify_config), "get_model_credits", return_value=3), + patch("services.credit_pool_service.CreditPoolService._use_billing_quota", return_value=True), + patch("services.billing_service.BillingService.quota_reserve", side_effect=reserve), + patch("services.billing_service.BillingService.quota_commit", side_effect=commit), + patch.object(ModelInstance, "invoke_llm", side_effect=provider_invoke), + ): + first = list(service.invoke(prepared)) + second = list(service.invoke(prepared)) + + assert first[0].delta.message.content == "done" + assert second[0].delta.message.content == "done" + assert reserve_request_ids == [request.caller.invocation_id, request.caller.invocation_id] + assert reservations_by_request == {request.caller.invocation_id: "reservation-1"} + assert committed_reservations == {"reservation-1"} + assert effective_charges == 3 + + def test_gateway_releases_reservation_when_provider_fails_before_delivery( sqlite_session_factory: sessionmaker[Session], sqlite_session: Session,