Merge remote-tracking branch 'origin/deploy/agent' into deploy/agent

This commit is contained in:
Joel 2026-08-13 16:08:06 +08:00
commit 7a40fc15d9
5 changed files with 146 additions and 11 deletions

View File

@ -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},
)

View File

@ -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(

View File

@ -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,

View File

@ -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()

View File

@ -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,