diff --git a/api/core/app/llm/__init__.py b/api/core/app/llm/__init__.py index d20a5b2344d..6f4e6909e18 100644 --- a/api/core/app/llm/__init__.py +++ b/api/core/app/llm/__init__.py @@ -5,6 +5,7 @@ from .quota import ( deduct_llm_quota_for_model, ensure_llm_quota_available, ensure_llm_quota_available_for_model, + reserve_llm_quota_for_model, ) __all__ = [ @@ -12,4 +13,5 @@ __all__ = [ "deduct_llm_quota_for_model", "ensure_llm_quota_available", "ensure_llm_quota_available_for_model", + "reserve_llm_quota_for_model", ] diff --git a/api/core/app/llm/quota.py b/api/core/app/llm/quota.py index d26d5d8a998..e4e84502cfa 100644 --- a/api/core/app/llm/quota.py +++ b/api/core/app/llm/quota.py @@ -7,6 +7,10 @@ with a non-LLM model. """ import warnings +from dataclasses import dataclass, field +from enum import StrEnum, auto +from typing import Any +from uuid import uuid4 from sqlalchemy import select from sqlalchemy.orm import sessionmaker @@ -23,6 +27,68 @@ from graphon.model_runtime.entities.model_entities import ModelType from libs.datetime_utils import naive_utc_now from models.provider import Provider, ProviderType from models.provider_ids import ModelProviderID +from services.credit_pool_service import CreditPoolReservation, CreditPoolService + + +class LLMQuotaReservationState(StrEnum): + RESERVED = auto() + COMMITTED = auto() + RELEASED = auto() + + +@dataclass +class LLMQuotaReservation: + """Quota reserved for one system-hosted LLM invocation.""" + + tenant_id: str + provider: str + model: str + provider_configuration: Any + quota_unit: QuotaUnit | None = None + credit_pool_reservation: CreditPoolReservation | None = None + requires_usage: bool = False + _state: LLMQuotaReservationState = field(default=LLMQuotaReservationState.RESERVED, init=False, repr=False) + + @property + def state(self) -> LLMQuotaReservationState: + return self._state + + @property + def commit_before_delivery(self) -> bool: + return self.credit_pool_reservation is not None + + def commit(self, usage: LLMUsage | None = None) -> None: + if self._state == LLMQuotaReservationState.COMMITTED: + return + if self._state == LLMQuotaReservationState.RELEASED: + raise RuntimeError("Cannot commit a released LLM quota reservation.") + + if self.credit_pool_reservation is not None: + self.credit_pool_reservation.commit() + elif self.requires_usage: + if usage is None: + raise ValueError("Accurate terminal usage is required for token-based LLM quota settlement.") + used_quota = _resolve_llm_used_quota( + system_configuration=self.provider_configuration.system_configuration, + model=self.model, + usage=usage, + ) + _deduct_used_llm_quota( + tenant_id=self.tenant_id, + provider=self.provider, + provider_configuration=self.provider_configuration, + used_quota=used_quota, + ) + + self._state = LLMQuotaReservationState.COMMITTED + + def release(self) -> None: + if self._state in {LLMQuotaReservationState.COMMITTED, LLMQuotaReservationState.RELEASED}: + return + + if self.credit_pool_reservation is not None: + self.credit_pool_reservation.release() + self._state = LLMQuotaReservationState.RELEASED def _get_provider_configuration(*, tenant_id: str, provider: str): @@ -34,6 +100,67 @@ def _get_provider_configuration(*, tenant_id: str, provider: str): return provider_configuration +def _get_current_quota_configuration(system_configuration): + return next( + ( + quota_configuration + for quota_configuration in system_configuration.quota_configurations + if quota_configuration.quota_type == system_configuration.current_quota_type + ), + None, + ) + + +def reserve_llm_quota_for_model(*, tenant_id: str, provider: str, model: str) -> LLMQuotaReservation: + """Reserve system-hosted LLM quota before invoking the provider.""" + provider_configuration = _get_provider_configuration(tenant_id=tenant_id, provider=provider) + reservation = LLMQuotaReservation( + tenant_id=tenant_id, + provider=provider, + model=model, + provider_configuration=provider_configuration, + ) + if provider_configuration.using_provider_type != ProviderType.SYSTEM: + return reservation + + provider_model = provider_configuration.get_provider_model(model_type=ModelType.LLM, model=model) + if provider_model and provider_model.status == ModelStatus.QUOTA_EXCEEDED: + raise QuotaExceededError(f"Model provider {provider} quota exceeded.") + + system_configuration = provider_configuration.system_configuration + quota_configuration = _get_current_quota_configuration(system_configuration) + if quota_configuration is None or quota_configuration.quota_limit == -1: + return reservation + + reservation.quota_unit = quota_configuration.quota_unit + quota_type = system_configuration.current_quota_type + if quota_type in {ProviderQuotaType.TRIAL, ProviderQuotaType.PAID}: + match quota_configuration.quota_unit: + case QuotaUnit.CREDITS: + amount = dify_config.get_model_credits(model) + case QuotaUnit.TIMES: + amount = 1 + case QuotaUnit.TOKENS: + # Token usage is unknown before invocation. Enabling TOKENS for a hosted + # credit pool requires accurate terminal usage and an upper-bound reservation strategy. + raise ValueError("Token-based hosted credit pools do not support pre-invocation reservation.") + case _: + raise ValueError(f"Unsupported hosted credit pool quota unit: {quota_configuration.quota_unit}") + + reservation.credit_pool_reservation = CreditPoolService.reserve_credits( + tenant_id=tenant_id, + credits_required=amount, + pool_type="paid" if quota_type == ProviderQuotaType.PAID else "trial", + request_id=str(uuid4()), + session_factory=db.session, + meta={"source": "llm.invoke", "provider": provider, "model": model}, + ) + elif quota_type == ProviderQuotaType.FREE: + reservation.requires_usage = True + + return reservation + + def ensure_llm_quota_available_for_model(*, tenant_id: str, provider: str, model: str) -> None: """Raise when a tenant-bound LLM model is already out of quota.""" provider_configuration = _get_provider_configuration(tenant_id=tenant_id, provider=provider) diff --git a/api/core/app/workflow/layers/__init__.py b/api/core/app/workflow/layers/__init__.py index 7d5841275db..945f75303c7 100644 --- a/api/core/app/workflow/layers/__init__.py +++ b/api/core/app/workflow/layers/__init__.py @@ -1,11 +1,9 @@ """Workflow-level GraphEngine layers that depend on outer infrastructure.""" -from .llm_quota import LLMQuotaLayer from .observability import ObservabilityLayer from .persistence import PersistenceWorkflowInfo, WorkflowPersistenceLayer __all__ = [ - "LLMQuotaLayer", "ObservabilityLayer", "PersistenceWorkflowInfo", "WorkflowPersistenceLayer", diff --git a/api/core/app/workflow/layers/llm_quota.py b/api/core/app/workflow/layers/llm_quota.py deleted file mode 100644 index 2422eed5a70..00000000000 --- a/api/core/app/workflow/layers/llm_quota.py +++ /dev/null @@ -1,194 +0,0 @@ -""" -LLM quota deduction layer for GraphEngine. - -This layer centralizes model-quota handling outside node implementations. - -Graphon LLM-backed nodes expose provider/model identity through public node -configuration and, after execution, through ``node_run_result.inputs``. Resolve -quota billing from that public identity instead of depending on -``ModelInstance`` reconstruction inside the workflow layer. Missing identity on -quota-tracked nodes is treated as a workflow bug and aborts execution so quota -handling is never silently skipped. -""" - -import logging -from typing import final, override - -from core.app.llm import deduct_llm_quota_for_model, ensure_llm_quota_available_for_model -from core.errors.error import QuotaExceededError -from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionStatus -from graphon.graph_engine.entities.commands import AbortCommand, CommandType -from graphon.graph_engine.layers import GraphEngineLayer -from graphon.graph_events import GraphEngineEvent, GraphNodeEventBase, NodeRunSucceededEvent -from graphon.node_events import NodeRunResult -from graphon.nodes.base.node import Node - -logger = logging.getLogger(__name__) -_QUOTA_NODE_TYPES = frozenset( - [ - BuiltinNodeTypes.LLM, - BuiltinNodeTypes.PARAMETER_EXTRACTOR, - BuiltinNodeTypes.QUESTION_CLASSIFIER, - ] -) - - -@final -class LLMQuotaLayer(GraphEngineLayer): - """Graph layer that applies tenant-scoped quota checks to LLM-backed nodes.""" - - tenant_id: str - _abort_sent: bool - - def __init__(self, tenant_id: str) -> None: - super().__init__() - self.tenant_id = tenant_id - self._abort_sent = False - - @override - def on_graph_start(self) -> None: - self._abort_sent = False - - @override - def on_event(self, event: GraphEngineEvent) -> None: - _ = event - - @override - def on_graph_end(self, error: Exception | None) -> None: - _ = error - - @override - def on_node_run_start(self, node: Node) -> None: - if self._abort_sent: - return - - if not self._supports_quota(node): - return - - model_identity = self._extract_model_identity_from_node(node) - if model_identity is None: - reason = "LLM quota check requires public node model identity before execution." - self._abort_before_node_run(node=node, reason=reason, error_type="LLMQuotaIdentityError") - logger.error("LLM quota handling aborted, node_id=%s, reason=%s", node.id, reason) - return - - provider, model_name = model_identity - try: - ensure_llm_quota_available_for_model( - tenant_id=self.tenant_id, - provider=provider, - model=model_name, - ) - except QuotaExceededError as exc: - self._abort_before_node_run(node=node, reason=str(exc), error_type=QuotaExceededError.__name__) - logger.warning("LLM quota check failed, node_id=%s, error=%s", node.id, exc) - - @override - def on_node_run_end( - self, node: Node, error: Exception | None, result_event: GraphNodeEventBase | None = None - ) -> None: - if error is not None or not isinstance(result_event, NodeRunSucceededEvent) or not self._supports_quota(node): - return - - model_identity = self._extract_model_identity_from_result_event(result_event) - if model_identity is None: - self._abort_for_missing_model_identity( - node=node, - reason="LLM quota deduction requires model identity in the node result event.", - ) - return - - provider, model_name = model_identity - - try: - deduct_llm_quota_for_model( - tenant_id=self.tenant_id, - provider=provider, - model=model_name, - usage=result_event.node_run_result.llm_usage, - ) - except QuotaExceededError as exc: - self._set_stop_event(node) - self._send_abort_command(reason=str(exc)) - logger.warning("LLM quota deduction exceeded, node_id=%s, error=%s", node.id, exc) - except Exception: - logger.exception("LLM quota deduction failed, node_id=%s", node.id) - - @staticmethod - def _set_stop_event(node: Node) -> None: - stop_event = getattr(node.graph_runtime_state, "stop_event", None) - if stop_event is not None: - stop_event.set() - - def _abort_before_node_run(self, *, node: Node, reason: str, error_type: str) -> None: - self._set_stop_event(node) - node.node_data.error_strategy = None - node.node_data.retry_config.retry_enabled = False - - def quota_aborted_run() -> NodeRunResult: - return NodeRunResult( - status=WorkflowNodeExecutionStatus.FAILED, - error=reason, - error_type=error_type, - ) - - # TODO: Push Graphon to expose a public pre-run failure/skip hook, then replace this private _run override. - node._run = quota_aborted_run # type: ignore[method-assign] - self._send_abort_command(reason=reason) - - def _abort_for_missing_model_identity(self, *, node: Node, reason: str) -> None: - self._set_stop_event(node) - self._send_abort_command(reason=reason) - logger.error("LLM quota handling aborted, node_id=%s, reason=%s", node.id, reason) - - def _send_abort_command(self, *, reason: str) -> None: - if not self.command_channel or self._abort_sent: - return - - try: - self.command_channel.send_command( - AbortCommand( - command_type=CommandType.ABORT, - reason=reason, - ) - ) - self._abort_sent = True - except Exception: - logger.exception("Failed to send quota abort command") - - @staticmethod - def _supports_quota(node: Node) -> bool: - return node.node_type in _QUOTA_NODE_TYPES - - @staticmethod - def _extract_model_identity_from_result_event(result_event: NodeRunSucceededEvent) -> tuple[str, str] | None: - provider = result_event.node_run_result.inputs.get("model_provider") - model_name = result_event.node_run_result.inputs.get("model_name") - if isinstance(provider, str) and provider and isinstance(model_name, str) and model_name: - return provider, model_name - return None - - @staticmethod - def _extract_model_identity_from_node(node: Node) -> tuple[str, str] | None: - node_data = getattr(node, "node_data", None) - if node_data is None: - node_data = getattr(node, "data", None) - - model_config = getattr(node_data, "model", None) - if model_config is None: - logger.warning( - "LLMQuotaLayer skipped quota handling because node model config is missing, node_id=%s", - node.id, - ) - return None - - provider = getattr(model_config, "provider", None) - model_name = getattr(model_config, "name", None) - if isinstance(provider, str) and provider and isinstance(model_name, str) and model_name: - return provider, model_name - - logger.warning( - "LLMQuotaLayer skipped quota handling because node model identity is invalid, node_id=%s", - node.id, - ) - return None diff --git a/api/core/entities/provider_configuration.py b/api/core/entities/provider_configuration.py index e933165b316..a73160ed1de 100644 --- a/api/core/entities/provider_configuration.py +++ b/api/core/entities/provider_configuration.py @@ -1857,10 +1857,10 @@ class ProviderConfiguration(BaseModel): ) ) - # if llm name not in restricted llm list, remove it + # Hosted allowlists currently use exact model names across model types. restrict_model_names = [rm.model for rm in restrict_models] for provider_model in provider_models: - if provider_model.model_type == ModelType.LLM and provider_model.model not in restrict_model_names: + if provider_model.model not in restrict_model_names: provider_model.status = ModelStatus.NO_PERMISSION elif not quota_configuration.is_valid: provider_model.status = ModelStatus.QUOTA_EXCEEDED diff --git a/api/core/model_manager.py b/api/core/model_manager.py index 29113ac6b2c..c07cc74583d 100644 --- a/api/core/model_manager.py +++ b/api/core/model_manager.py @@ -1,19 +1,19 @@ 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 +from typing import IO, Any, Literal, Optional, ParamSpec, TypeVar, Union, cast, overload, override from configs import dify_config from core.entities import PluginCredentialType from core.entities.embedding_type import EmbeddingInputType from core.entities.provider_configuration import ProviderConfiguration, ProviderModelBundle from core.entities.provider_entities import ModelLoadBalancingConfiguration -from core.errors.error import ProviderTokenNotInitError +from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError from core.plugin.impl.model_runtime_factory import create_plugin_provider_manager from core.provider_manager import ProviderManager from extensions.ext_redis import redis_client from graphon.model_runtime.callbacks.base_callback import Callback -from graphon.model_runtime.entities.llm_entities import LLMResult +from graphon.model_runtime.entities.llm_entities import LLMResult, LLMUsage from graphon.model_runtime.entities.message_entities import PromptMessage, PromptMessageTool from graphon.model_runtime.entities.model_entities import AIModelEntity, ModelFeature, ModelType from graphon.model_runtime.entities.rerank_entities import MultimodalRerankInput, RerankResult @@ -442,6 +442,149 @@ class ModelInstance: ) +class QuotaManagedModelInstance(ModelInstance): + """A system-hosted LLM instance that owns quota settlement per invocation.""" + + def reserve_quota(self): + 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, + ) + + @staticmethod + def release_quota_safely(reservation) -> None: + try: + reservation.release() + except Exception: + logger.exception("Failed to release LLM quota reservation") + + @overload + def invoke_llm( + self, + prompt_messages: Sequence[PromptMessage], + model_parameters: dict[str, Any] | None = None, + tools: Sequence[PromptMessageTool] | None = None, + stop: list[str] | None = None, + stream: Literal[True] = True, + callbacks: list[Callback] | None = None, + request_metadata: Mapping[str, object] | None = None, + ) -> Generator: ... + + @overload + def invoke_llm( + self, + prompt_messages: list[PromptMessage], + model_parameters: dict[str, Any] | None = None, + tools: Sequence[PromptMessageTool] | None = None, + stop: list[str] | None = None, + stream: Literal[False] = False, + callbacks: list[Callback] | None = None, + request_metadata: Mapping[str, object] | None = None, + ) -> LLMResult: ... + + @overload + def invoke_llm( + self, + prompt_messages: list[PromptMessage], + model_parameters: dict[str, Any] | None = None, + tools: Sequence[PromptMessageTool] | None = None, + stop: list[str] | None = None, + stream: bool = True, + callbacks: list[Callback] | None = None, + request_metadata: Mapping[str, object] | None = None, + ) -> Union[LLMResult, Generator]: ... + + @override + def invoke_llm( + self, + prompt_messages: Sequence[PromptMessage], + model_parameters: dict[str, Any] | None = None, + tools: Sequence[PromptMessageTool] | None = None, + stop: Sequence[str] | None = None, + stream: bool = True, + callbacks: list[Callback] | None = None, + request_metadata: Mapping[str, object] | None = None, + ) -> Union[LLMResult, Generator]: + normalized_prompt_messages = list(prompt_messages) + normalized_stop = list(stop) if stop else None + if stream: + return self._invoke_llm_stream( + prompt_messages=normalized_prompt_messages, + model_parameters=model_parameters, + tools=tools, + stop=normalized_stop, + callbacks=callbacks, + request_metadata=request_metadata, + ) + + reservation = self.reserve_quota() + try: + response = super().invoke_llm( + prompt_messages=normalized_prompt_messages, + model_parameters=model_parameters, + tools=tools, + stop=normalized_stop, + stream=False, + callbacks=callbacks, + request_metadata=request_metadata, + ) + if isinstance(response, Generator): + raise TypeError("Non-streaming LLM invocation returned a generator.") + reservation.commit(response.usage) + return response + finally: + self.release_quota_safely(reservation) + + def _invoke_llm_stream( + self, + *, + prompt_messages: list[PromptMessage], + model_parameters: dict[str, Any] | None, + tools: Sequence[PromptMessageTool] | None, + stop: list[str] | None, + callbacks: list[Callback] | None, + request_metadata: Mapping[str, object] | None, + ) -> Generator: + reservation = self.reserve_quota() + usage: LLMUsage | None = None + try: + response = super().invoke_llm( + prompt_messages=prompt_messages, + model_parameters=model_parameters, + tools=tools, + stop=stop, + stream=True, + callbacks=callbacks, + request_metadata=request_metadata, + ) + if not isinstance(response, Generator): + raise TypeError("Streaming LLM invocation did not return a generator.") + + if reservation.commit_before_delivery: + for chunk in response: + chunk_usage = chunk.delta.usage + if chunk_usage is not None: + usage = chunk_usage + reservation.commit(usage) + yield chunk + return + + buffered_chunks = [] + for chunk in response: + chunk_usage = chunk.delta.usage + if chunk_usage is not None: + usage = chunk_usage + buffered_chunks.append(chunk) + + reservation.commit(usage) + yield from buffered_chunks + finally: + self.release_quota_safely(reservation) + + class ModelManager: """Resolves :class:`ModelInstance` objects for a tenant and provider. @@ -472,6 +615,43 @@ class ModelManager: def for_tenant(cls, tenant_id: str, user_id: str | None = None) -> "ModelManager": return cls(provider_manager=create_plugin_provider_manager(tenant_id=tenant_id, user_id=user_id)) + @staticmethod + def _validate_system_model_access( + provider_model_bundle: ProviderModelBundle, + *, + model_type: ModelType, + model: str, + ) -> None: + configuration = provider_model_bundle.configuration + if configuration.using_provider_type != ProviderType.SYSTEM: + return + + # Hosted allowlists retain the existing comma-separated format. Model names + # are matched exactly; model-type-specific entries will be introduced later. + quota_configuration = next( + ( + quota + for quota in configuration.system_configuration.quota_configurations + if quota.quota_type == configuration.system_configuration.current_quota_type + ), + None, + ) + if quota_configuration is None or not quota_configuration.restrict_models: + return + if any(restricted_model.model == model for restricted_model in quota_configuration.restrict_models): + return + + raise ModelCurrentlyNotSupportError(f"System model {model_type.value}/{model} is not allowed.") + + @staticmethod + def _model_instance_class(provider_model_bundle: ProviderModelBundle, model_type: ModelType) -> type[ModelInstance]: + if ( + model_type == ModelType.LLM + and provider_model_bundle.configuration.using_provider_type == ProviderType.SYSTEM + ): + return QuotaManagedModelInstance + return ModelInstance + def get_model_instance( self, tenant_id: str, @@ -493,17 +673,19 @@ class ModelManager: provider_model_bundle = self._provider_manager.get_provider_model_bundle( tenant_id=tenant_id, provider=provider, model_type=model_type ) + self._validate_system_model_access(provider_model_bundle, model_type=model_type, model=model) + model_instance_class = self._model_instance_class(provider_model_bundle, model_type) cred_cache_key = (tenant_id, provider, model_type.value, model) if cred_cache_key in self._credentials_cache: - return ModelInstance( + return model_instance_class( provider_model_bundle, model, deepcopy(self._credentials_cache[cred_cache_key]), ) - ret = ModelInstance(provider_model_bundle, model) + ret = model_instance_class(provider_model_bundle, model) if self._enable_credentials_cache: self._credentials_cache[cred_cache_key] = deepcopy(ret.credentials) return ret diff --git a/api/core/plugin/backwards_invocation/model.py b/api/core/plugin/backwards_invocation/model.py index c03665272b9..df2dd7c795f 100644 --- a/api/core/plugin/backwards_invocation/model.py +++ b/api/core/plugin/backwards_invocation/model.py @@ -3,7 +3,6 @@ from binascii import hexlify, unhexlify from collections.abc import Generator from typing import Any -from core.app.llm import deduct_llm_quota from core.llm_generator.output_parser.structured_output import invoke_llm_with_structured_output from core.model_manager import ModelManager from core.plugin.backwards_invocation.base import BaseBackwardsInvocation @@ -80,15 +79,11 @@ class PluginModelBackwardsInvocation(BaseBackwardsInvocation): def handle() -> Generator[LLMResultChunk, None, None]: for chunk in response: - if chunk.delta.usage: - deduct_llm_quota(tenant_id=tenant.id, model_instance=model_instance, usage=chunk.delta.usage) chunk.prompt_messages = [] yield chunk return handle() else: - if response.usage: - deduct_llm_quota(tenant_id=tenant.id, model_instance=model_instance, usage=response.usage) def handle_non_streaming(response: LLMResult) -> Generator[LLMResultChunk, None, None]: yield LLMResultChunk( @@ -141,15 +136,11 @@ class PluginModelBackwardsInvocation(BaseBackwardsInvocation): def handle() -> Generator[LLMResultChunkWithStructuredOutput, None, None]: for chunk in response: - if chunk.delta.usage: - deduct_llm_quota(tenant_id=tenant.id, model_instance=model_instance, usage=chunk.delta.usage) chunk.prompt_messages = [] yield chunk return handle() else: - if response.usage: - deduct_llm_quota(tenant_id=tenant.id, model_instance=model_instance, usage=response.usage) def handle_non_streaming( response: LLMResultWithStructuredOutput, diff --git a/api/core/rag/retrieval/router/multi_dataset_react_route.py b/api/core/rag/retrieval/router/multi_dataset_react_route.py index 21a9d04f7f2..95ffd4c84ca 100644 --- a/api/core/rag/retrieval/router/multi_dataset_react_route.py +++ b/api/core/rag/retrieval/router/multi_dataset_react_route.py @@ -2,7 +2,6 @@ from collections.abc import Generator, Sequence from typing import Any, Union from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity -from core.app.llm import deduct_llm_quota from core.model_manager import ModelInstance, ModelManager from core.prompt.advanced_prompt_transform import AdvancedPromptTransform from core.prompt.entities.advanced_prompt_entities import ChatModelMessage, CompletionModelPromptTemplate @@ -168,9 +167,6 @@ class ReactMultiDatasetRouter: # handle invoke result text, usage = self._handle_invoke_result(invoke_result=invoke_result) - # deduct quota - deduct_llm_quota(tenant_id=tenant_id, model_instance=bound_model_instance, usage=usage) - return text, usage def _handle_invoke_result(self, invoke_result: Generator) -> tuple[str, LLMUsage]: diff --git a/api/core/workflow/llm_node.py b/api/core/workflow/llm_node.py new file mode 100644 index 00000000000..1f1c863ab5a --- /dev/null +++ b/api/core/workflow/llm_node.py @@ -0,0 +1,44 @@ +from collections.abc import Callable, Generator, Sequence +from typing import Any, override + +from graphon.model_runtime.entities.llm_entities import LLMStructuredOutput +from graphon.model_runtime.entities.message_entities import PromptMessage +from graphon.node_events.base import NodeEventBase +from graphon.nodes.llm.node import LLMNode +from graphon.nodes.llm.runtime_protocols import LLMPollingCapableProtocol + + +# TODO: Remove this Dify-specific node once graphon exposes a polling finalization hook. +class DifyLLMNode(LLMNode): + """Dify-owned LLM node lifecycle extensions.""" + + @classmethod + @override + def version(cls) -> str: + return "1" + + def __init__( + self, + *args: Any, + polling_finalizer: Callable[[], None], + **kwargs: Any, + ) -> None: + super().__init__(*args, **kwargs) + self._polling_finalizer = polling_finalizer + + @override + def _invoke_llm_with_polling( + self, + *, + polling_model: LLMPollingCapableProtocol, + prompt_messages: Sequence[PromptMessage], + stop: Sequence[str] | None, + ) -> Generator[NodeEventBase | LLMStructuredOutput, None, None]: + try: + yield from super()._invoke_llm_with_polling( + polling_model=polling_model, + prompt_messages=prompt_messages, + stop=stop, + ) + finally: + self._polling_finalizer() diff --git a/api/core/workflow/node_factory.py b/api/core/workflow/node_factory.py index 4807a945e61..f73dd47e7bc 100644 --- a/api/core/workflow/node_factory.py +++ b/api/core/workflow/node_factory.py @@ -27,6 +27,7 @@ from core.workflow.llm_environment_variable import ( resolve_llm_model_config, should_resolve_llm_model_selector, ) +from core.workflow.llm_node import DifyLLMNode from core.workflow.node_runtime import ( DifyFileReferenceFactory, DifyHumanInputNodeRuntime, @@ -493,6 +494,8 @@ class DifyNodeFactory(NodeFactory): @staticmethod def _resolve_node_class(*, node_type: NodeType, node_version: str) -> type[Node]: + if node_type == BuiltinNodeTypes.LLM: + return DifyLLMNode return resolve_workflow_node_class(node_type=node_type, node_version=node_version) def _resolve_llm_model_reference(self, node_data: LLMNodeData) -> LLMNodeData: @@ -586,18 +589,19 @@ class DifyNodeFactory(NodeFactory): ) -> dict[str, object]: validated_node_data = cast(LLMCompatibleNodeData, node_data) model_instance = self._build_model_instance_for_llm_node(validated_node_data) + node_model_instance = ( + self._wrap_model_instance_for_node( + node_data=validated_node_data, + model_instance=model_instance, + request_metadata={"app_id": self._dify_context.app_id}, + ) + if wrap_model_instance + else model_instance + ) node_init_kwargs: dict[str, object] = { "credentials_provider": self._llm_credentials_provider, "model_factory": self._llm_model_factory, - "model_instance": ( - self._wrap_model_instance_for_node( - node_data=validated_node_data, - model_instance=model_instance, - request_metadata={"app_id": self._dify_context.app_id}, - ) - if wrap_model_instance - else model_instance - ), + "model_instance": node_model_instance, "memory": self._build_memory_for_llm_node( node_data=validated_node_data, model_instance=model_instance, @@ -619,6 +623,7 @@ class DifyNodeFactory(NodeFactory): node_init_kwargs["jinja2_template_renderer"] = self._jinja2_template_renderer if validated_node_data.type == BuiltinNodeTypes.LLM: node_init_kwargs["default_query_selector"] = system_variable_selector(SystemVariableKey.QUERY) + node_init_kwargs["polling_finalizer"] = cast(DifyPreparedLLM, node_model_instance).finalize_llm_polling return node_init_kwargs @staticmethod diff --git a/api/core/workflow/node_runtime.py b/api/core/workflow/node_runtime.py index d0391b69c73..f113a88ccb9 100644 --- a/api/core/workflow/node_runtime.py +++ b/api/core/workflow/node_runtime.py @@ -3,7 +3,7 @@ from __future__ import annotations from collections.abc import Callable, Generator, Mapping, Sequence from dataclasses import dataclass from enum import Enum -from typing import TYPE_CHECKING, Any, Literal, cast, overload, override +from typing import TYPE_CHECKING, Any, Literal, Protocol, cast, overload, override from pydantic import JsonValue from sqlalchemy import select @@ -20,7 +20,7 @@ from core.db.session_factory import session_factory from core.helper.trace_id_helper import ParentTraceContext from core.llm_generator.output_parser.errors import OutputParserError from core.llm_generator.output_parser.structured_output import invoke_llm_with_structured_output -from core.model_manager import ModelInstance +from core.model_manager import ModelInstance, QuotaManagedModelInstance from core.plugin.impl.exc import PluginDaemonClientSideError, PluginInvokeError from core.plugin.impl.plugin import PluginInstaller from core.prompt.utils.prompt_message_util import PromptMessageUtil @@ -49,6 +49,7 @@ from graphon.file import File, FileTransferMethod, FileType from graphon.model_runtime.entities import LLMMode from graphon.model_runtime.entities.llm_entities import ( LLMPollingResult, + LLMPollingStatus, LLMResult, LLMResultChunk, LLMResultChunkWithStructuredOutput, @@ -87,6 +88,33 @@ from .human_input_adapter import ( ) from .system_variables import SystemVariableKey, get_system_text + +class PollingLLMRuntimeProtocol(Protocol): + """Runtime capability required by the workflow polling adapter.""" + + def start_llm_polling( + self, + *, + provider: str, + model: str, + credentials: dict[str, Any], + model_parameters: dict[str, Any], + prompt_messages: Sequence[PromptMessage], + tools: Sequence[PromptMessageTool] | None, + stop: Sequence[str] | None, + json_schema: dict[str, Any] | None, + ) -> LLMPollingResult: ... + + def check_llm_polling( + self, + *, + provider: str, + model: str, + credentials: dict[str, Any], + plugin_state: dict[str, JsonValue], + ) -> LLMPollingResult: ... + + if TYPE_CHECKING: from core.tools.__base.tool import Tool from core.tools.entities.tool_entities import ToolInvokeMessage as CoreToolInvokeMessage @@ -281,23 +309,45 @@ class DifyPreparedLLM(LLMProtocol): def is_structured_output_parse_error(self, error: Exception) -> bool: return isinstance(error, OutputParserError) + def finalize_llm_polling(self) -> None: + """Finalize resources held by a polling invocation, if any.""" + class DifyPreparedPollingLLM(DifyPreparedLLM, LLMPollingCapableProtocol): """Prepared workflow LLM adapter that exposes Graphon's polling protocol.""" def __init__(self, model_instance: ModelInstance, request_metadata: Mapping[str, object] | None = None) -> None: - from core.plugin.impl.model_runtime import PluginModelRuntime - super().__init__(model_instance, request_metadata=request_metadata) - model_type_instance = model_instance.model_type_instance - if not isinstance(model_type_instance, LargeLanguageModel): - raise TypeError("Polling wrapper requires a large-language-model instance.") + model_type_instance = cast(LargeLanguageModel, model_instance.model_type_instance) + self._polling_runtime = cast(PollingLLMRuntimeProtocol, model_type_instance.model_runtime) + self._polling_quota_reservation = None - plugin_model_runtime = model_type_instance.model_runtime - if not isinstance(plugin_model_runtime, PluginModelRuntime): - raise TypeError("Polling wrapper requires a plugin-backed model runtime.") + @override + def finalize_llm_polling(self) -> None: + reservation = self._polling_quota_reservation + self._polling_quota_reservation = None + if reservation is not None: + QuotaManagedModelInstance.release_quota_safely(reservation) - self._plugin_model_runtime = plugin_model_runtime + def _settle_polling_quota(self, polling_result: LLMPollingResult) -> LLMPollingResult: + reservation = self._polling_quota_reservation + if reservation is None or polling_result.status == LLMPollingStatus.RUNNING: + return polling_result + + try: + if polling_result.status == LLMPollingStatus.SUCCEEDED: + if polling_result.result is None: + raise ValueError("A successful LLM polling result must include a model result.") + reservation.commit(polling_result.result.usage) + else: + reservation.release() + except Exception: + QuotaManagedModelInstance.release_quota_safely(reservation) + raise + finally: + self._polling_quota_reservation = None + + return polling_result @override def start_llm_polling( @@ -309,16 +359,26 @@ class DifyPreparedPollingLLM(DifyPreparedLLM, LLMPollingCapableProtocol): stop: Sequence[str] | None, json_schema: Mapping[str, Any] | None, ) -> LLMPollingResult: - return self._plugin_model_runtime.start_llm_polling( - provider=self.provider, - model=self.model_name, - credentials=self._model_instance.credentials, - prompt_messages=prompt_messages, - model_parameters=dict(model_parameters), - tools=tools, - stop=stop, - json_schema=dict(json_schema) if json_schema is not None else None, - ) + self.finalize_llm_polling() + + if isinstance(self._model_instance, QuotaManagedModelInstance): + self._polling_quota_reservation = self._model_instance.reserve_quota() + + try: + polling_result = self._polling_runtime.start_llm_polling( + provider=self.provider, + model=self.model_name, + credentials=self._model_instance.credentials, + prompt_messages=prompt_messages, + model_parameters=dict(model_parameters), + tools=tools, + stop=stop, + json_schema=dict(json_schema) if json_schema is not None else None, + ) + return self._settle_polling_quota(polling_result) + except Exception: + self.finalize_llm_polling() + raise @override def check_llm_polling( @@ -326,12 +386,17 @@ class DifyPreparedPollingLLM(DifyPreparedLLM, LLMPollingCapableProtocol): *, plugin_state: Mapping[str, JsonValue], ) -> LLMPollingResult: - return self._plugin_model_runtime.check_llm_polling( - provider=self.provider, - model=self.model_name, - credentials=self._model_instance.credentials, - plugin_state=dict(plugin_state), - ) + try: + polling_result = self._polling_runtime.check_llm_polling( + provider=self.provider, + model=self.model_name, + credentials=self._model_instance.credentials, + plugin_state=dict(plugin_state), + ) + return self._settle_polling_quota(polling_result) + except Exception: + self.finalize_llm_polling() + raise class DifyPromptMessageSerializer(PromptMessageSerializerProtocol): diff --git a/api/core/workflow/workflow_entry.py b/api/core/workflow/workflow_entry.py index d74fade25e1..866bc73fcf6 100644 --- a/api/core/workflow/workflow_entry.py +++ b/api/core/workflow/workflow_entry.py @@ -9,7 +9,6 @@ from context import capture_current_context from core.app.apps.exc import GenerateTaskStoppedError from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom, build_dify_run_context from core.app.file_access import DatabaseFileAccessController -from core.app.workflow.layers.llm_quota import LLMQuotaLayer from core.app.workflow.layers.observability import ObservabilityLayer from core.workflow.node_factory import ( DifyGraphInitContext, @@ -170,7 +169,6 @@ class WorkflowEntry: max_steps=dify_config.WORKFLOW_MAX_EXECUTION_STEPS, max_time=dify_config.WORKFLOW_MAX_EXECUTION_TIME ) self.graph_engine.layer(limits_layer) - self.graph_engine.layer(LLMQuotaLayer(tenant_id=tenant_id)) # Add observability layer when OTel is enabled if dify_config.ENABLE_OTEL or is_instrument_flag_enabled(): @@ -550,10 +548,7 @@ class WorkflowEntry: """ Run a standalone node with the same quota and observability hooks as GraphEngine. """ - layers: Sequence[GraphEngineLayer] = ( - LLMQuotaLayer(tenant_id=tenant_id), - ObservabilityLayer(), - ) + layers: Sequence[GraphEngineLayer] = (ObservabilityLayer(),) command_channel = InMemoryChannel() runtime_state = ReadOnlyGraphRuntimeStateWrapper(node.graph_runtime_state) for layer in layers: diff --git a/api/services/credit_pool_service.py b/api/services/credit_pool_service.py index deea8245ffc..95f4d90bc79 100644 --- a/api/services/credit_pool_service.py +++ b/api/services/credit_pool_service.py @@ -7,7 +7,9 @@ from piling up database transactions while preserving cross-tenant concurrency. import logging from collections.abc import Callable, Mapping -from dataclasses import dataclass +from dataclasses import dataclass, field +from enum import StrEnum, auto +from typing import Any from uuid import uuid4 from sqlalchemy import select @@ -44,6 +46,77 @@ class CreditPoolBalance: return self.quota_limit == -1 or self.remaining_credits >= required_credits +class CreditPoolReservationState(StrEnum): + RESERVED = auto() + COMMITTED = auto() + RELEASED = auto() + + +@dataclass +class CreditPoolReservation: + """A strict credit-pool reservation spanning one billable operation.""" + + tenant_id: str + pool_type: str + amount: int + request_id: str + reservation_id: str | None + meta: dict[str, Any] = field(default_factory=dict) + _session_factory: Callable[[], Session] | None = field(default=None, repr=False) + _state: CreditPoolReservationState = field(default=CreditPoolReservationState.RESERVED, init=False, repr=False) + + @property + def state(self) -> CreditPoolReservationState: + return self._state + + def commit(self) -> None: + if self._state == CreditPoolReservationState.COMMITTED: + return + if self._state == CreditPoolReservationState.RELEASED: + raise RuntimeError("Cannot commit a released credit reservation.") + + if self.reservation_id is not None: + from services.billing_service import BillingService + + BillingService.quota_commit( + tenant_id=self.tenant_id, + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket=self.pool_type, + reservation_id=self.reservation_id, + actual_amount=self.amount, + meta={**self.meta, "request_id": self.request_id}, + ) + + # The database fallback reserves by deducting under the tenant lock, so + # commit only makes that already durable reservation final. + self._state = CreditPoolReservationState.COMMITTED + + def release(self) -> None: + if self._state in {CreditPoolReservationState.COMMITTED, CreditPoolReservationState.RELEASED}: + return + + if self.reservation_id is not None: + from services.billing_service import BillingService + + BillingService.quota_release( + tenant_id=self.tenant_id, + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket=self.pool_type, + reservation_id=self.reservation_id, + ) + else: + if self._session_factory is None: + raise RuntimeError("Database credit reservation requires a session factory.") + CreditPoolService._release_database_reservation( + tenant_id=self.tenant_id, + pool_type=self.pool_type, + credits=self.amount, + session=self._session_factory(), + ) + + self._state = CreditPoolReservationState.RELEASED + + class CreditPoolService: @staticmethod def _normalize_pool_type(pool_type: str | ProviderQuotaType) -> str: @@ -162,6 +235,110 @@ class CreditPoolService: return False return pool.has_sufficient_credits(credits_required) + @classmethod + def reserve_credits( + cls, + tenant_id: str, + credits_required: int, + pool_type: str | ProviderQuotaType = "trial", + *, + request_id: str, + session_factory: Callable[[], Session] | None = None, + meta: dict[str, Any] | None = None, + ) -> CreditPoolReservation: + """Reserve the full amount or raise before the billable operation starts.""" + if credits_required <= 0: + raise ValueError("credits_required must be greater than 0") + if not request_id: + raise ValueError("request_id is required") + + normalized_pool_type = cls._normalize_pool_type(pool_type) + reservation_meta = {"source": "credit_pool.reservation", **(meta or {})} + if cls._use_billing_quota(): + from services.billing_service import BillingService + + result = BillingService.quota_reserve( + tenant_id=tenant_id, + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket=normalized_pool_type, + request_id=request_id, + amount=credits_required, + meta=reservation_meta, + ) + reservation_id = result.get("reservation_id", "") + if not reservation_id: + raise QuotaExceededError("Insufficient credits remaining") + return CreditPoolReservation( + tenant_id=tenant_id, + pool_type=normalized_pool_type, + amount=credits_required, + request_id=request_id, + reservation_id=reservation_id, + meta=reservation_meta, + ) + + if session_factory is None: + raise ValueError("session_factory is required when billing quota is disabled") + + session = session_factory() + + def reserve() -> int: + pool = cls._get_locked_pool(session=session, tenant_id=tenant_id, pool_type=normalized_pool_type) + if not pool: + raise QuotaExceededError("Credit pool not found") + if not pool.has_sufficient_credits(credits_required): + raise QuotaExceededError("Insufficient credits remaining") + + pool.quota_used += credits_required + session.commit() + return credits_required + + try: + cls._deduct_with_tenant_lock(tenant_id, reserve) + except QuotaExceededError: + session.rollback() + raise + except Exception: + session.rollback() + logger.exception("Failed to reserve credits for tenant %s", tenant_id) + raise QuotaExceededError("Failed to reserve credits") + + return CreditPoolReservation( + tenant_id=tenant_id, + pool_type=normalized_pool_type, + amount=credits_required, + request_id=request_id, + reservation_id=None, + meta=reservation_meta, + _session_factory=session_factory, + ) + + @classmethod + def _release_database_reservation( + cls, + *, + tenant_id: str, + pool_type: str, + credits: int, + session: Session, + ) -> None: + def release() -> int: + pool = cls._get_locked_pool(session=session, tenant_id=tenant_id, pool_type=pool_type) + if not pool: + raise QuotaExceededError("Credit pool not found") + if pool.quota_used < credits: + raise RuntimeError("Reserved credits exceed recorded usage.") + + pool.quota_used -= credits + session.commit() + return credits + + try: + cls._deduct_with_tenant_lock(tenant_id, release) + except Exception: + session.rollback() + raise + @classmethod def check_and_deduct_credits( cls, 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 ba691bb298a..e7297639610 100644 --- a/api/tests/unit_tests/core/app/test_llm_quota.py +++ b/api/tests/unit_tests/core/app/test_llm_quota.py @@ -10,10 +10,12 @@ from sqlalchemy.orm import sessionmaker from configs import dify_config from core.app.llm.quota import ( + LLMQuotaReservationState, deduct_llm_quota, deduct_llm_quota_for_model, ensure_llm_quota_available, ensure_llm_quota_available_for_model, + reserve_llm_quota_for_model, ) from core.entities.model_entities import ModelStatus from core.entities.provider_entities import ProviderQuotaType, QuotaUnit @@ -100,6 +102,111 @@ def test_ensure_llm_quota_available_for_model_ignores_custom_provider_configurat provider_configuration.get_provider_model.assert_not_called() +def test_reserve_llm_quota_uses_exact_credit_pool_reservation() -> 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.CREDITS, + 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.object(type(dify_config), "get_model_credits", return_value=9), + patch("core.app.llm.quota.CreditPoolService.reserve_credits", return_value=credit_reservation) as reserve, + ): + reservation = reserve_llm_quota_for_model( + tenant_id="tenant-id", + provider="openai", + model="gpt-4o", + ) + reservation.commit(LLMUsage.empty_usage()) + reservation.release() + + assert reservation.state == LLMQuotaReservationState.COMMITTED + assert reservation.commit_before_delivery is True + reserve.assert_called_once_with( + tenant_id="tenant-id", + credits_required=9, + pool_type="trial", + request_id=ANY, + session_factory=ANY, + meta={"source": "llm.invoke", "provider": "openai", "model": "gpt-4o"}, + ) + credit_reservation.commit.assert_called_once_with() + credit_reservation.release.assert_not_called() + + +def test_reserve_llm_quota_requires_accurate_usage_for_free_tokens() -> None: + 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.FREE, + quota_configurations=[ + SimpleNamespace( + quota_type=ProviderQuotaType.FREE, + quota_unit=QuotaUnit.TOKENS, + 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): + reservation = reserve_llm_quota_for_model( + tenant_id="tenant-id", + provider="openai", + model="gpt-4o", + ) + + assert reservation.commit_before_delivery is False + with pytest.raises(ValueError, match="Accurate terminal usage"): + reservation.commit() + + +def test_reserve_llm_quota_rejects_token_based_credit_pool() -> None: + 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.TOKENS, + 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), + pytest.raises(ValueError, match="do not support pre-invocation reservation"), + ): + reserve_llm_quota_for_model( + tenant_id="tenant-id", + provider="openai", + model="gpt-4o", + ) + + def test_deduct_llm_quota_for_model_uses_identity_based_trial_billing() -> None: usage = LLMUsage.empty_usage() usage.total_tokens = 42 diff --git a/api/tests/unit_tests/core/rag/retrieval/test_multi_dataset_react_route.py b/api/tests/unit_tests/core/rag/retrieval/test_multi_dataset_react_route.py index c56528cf55e..bcc8ee8fcb5 100644 --- a/api/tests/unit_tests/core/rag/retrieval/test_multi_dataset_react_route.py +++ b/api/tests/unit_tests/core/rag/retrieval/test_multi_dataset_react_route.py @@ -165,10 +165,7 @@ class TestReactMultiDatasetRouter: model_instance = Mock() model_instance.invoke_llm.return_value = iter([chunk]) - with ( - patch("core.rag.retrieval.router.multi_dataset_react_route.ModelManager.for_tenant") as mock_manager, - patch("core.rag.retrieval.router.multi_dataset_react_route.deduct_llm_quota") as mock_deduct, - ): + with patch("core.rag.retrieval.router.multi_dataset_react_route.ModelManager.for_tenant") as mock_manager: mock_manager.return_value.get_model_instance.return_value = model_instance text, returned_usage = router._invoke_llm( completion_param={"temperature": 0.1}, @@ -188,7 +185,6 @@ class TestReactMultiDatasetRouter: model_type=ModelType.LLM, model=model_instance.model_name, ) - mock_deduct.assert_called_once() def test_handle_invoke_result_with_empty_usage(self) -> None: router = ReactMultiDatasetRouter() diff --git a/api/tests/unit_tests/core/test_model_manager.py b/api/tests/unit_tests/core/test_model_manager.py index 5a7e7e30a50..fa873e8d865 100644 --- a/api/tests/unit_tests/core/test_model_manager.py +++ b/api/tests/unit_tests/core/test_model_manager.py @@ -4,10 +4,20 @@ import pytest import redis from pytest_mock import MockerFixture -from core.entities.provider_entities import ModelLoadBalancingConfiguration -from core.model_manager import LBModelManager, ModelManager +from core.entities.provider_entities import ( + ModelLoadBalancingConfiguration, + ProviderQuotaType, + QuotaConfiguration, + QuotaUnit, + RestrictModel, +) +from core.errors.error import ModelCurrentlyNotSupportError +from core.model_manager import LBModelManager, ModelInstance, ModelManager, QuotaManagedModelInstance from extensions.ext_redis import redis_client +from graphon.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk, LLMResultChunkDelta, LLMUsage +from graphon.model_runtime.entities.message_entities import AssistantPromptMessage from graphon.model_runtime.entities.model_entities import ModelType +from models.provider import ProviderType @pytest.fixture @@ -63,6 +73,208 @@ def test_model_manager_with_cache_enabled_reuses_stored_credentials(): get_creds.assert_called_once() +def _build_model_manager_bundle( + *, + provider_type: ProviderType, + restrict_models: list[RestrictModel], +) -> tuple[ModelManager, MagicMock]: + provider_manager = MagicMock() + bundle = MagicMock() + bundle.configuration.provider.provider = "openai" + bundle.configuration.tenant_id = "tenant-1" + bundle.configuration.model_settings = [] + bundle.configuration.using_provider_type = provider_type + bundle.configuration.system_configuration.current_quota_type = ProviderQuotaType.TRIAL + bundle.configuration.system_configuration.quota_configurations = [ + QuotaConfiguration( + quota_type=ProviderQuotaType.TRIAL, + quota_unit=QuotaUnit.CREDITS, + quota_limit=200, + quota_used=0, + is_valid=True, + restrict_models=restrict_models, + ) + ] + bundle.configuration.get_current_credentials.return_value = {"api_key": "hosted"} + bundle.model_type_instance.model_type = ModelType.LLM + provider_manager.get_provider_model_bundle.return_value = bundle + return ModelManager(provider_manager), bundle + + +def test_model_manager_wraps_allowlisted_system_llm() -> None: + manager, _ = _build_model_manager_bundle( + provider_type=ProviderType.SYSTEM, + restrict_models=[RestrictModel(model="gpt-4", model_type=ModelType.LLM)], + ) + + model_instance = manager.get_model_instance("tenant-1", "openai", ModelType.LLM, "gpt-4") + + assert isinstance(model_instance, QuotaManagedModelInstance) + + +def test_model_manager_rejects_system_model_by_exact_name() -> None: + manager, bundle = _build_model_manager_bundle( + provider_type=ProviderType.SYSTEM, + restrict_models=[RestrictModel(model="gpt-4", model_type=ModelType.LLM)], + ) + + with pytest.raises(ModelCurrentlyNotSupportError, match="llm/gpt-4o is not allowed"): + manager.get_model_instance("tenant-1", "openai", ModelType.LLM, "gpt-4o") + + bundle.configuration.get_current_credentials.assert_not_called() + + +def test_model_manager_matches_allowlist_name_across_model_types() -> None: + manager, _ = _build_model_manager_bundle( + provider_type=ProviderType.SYSTEM, + restrict_models=[RestrictModel(model="shared-model", model_type=ModelType.TEXT_EMBEDDING)], + ) + + model_instance = manager.get_model_instance("tenant-1", "openai", ModelType.LLM, "shared-model") + + assert isinstance(model_instance, QuotaManagedModelInstance) + + +def test_quota_managed_non_streaming_invocation_finalizes_reservation() -> None: + manager, _ = _build_model_manager_bundle( + provider_type=ProviderType.SYSTEM, + restrict_models=[RestrictModel(model="gpt-4", model_type=ModelType.LLM)], + ) + model_instance = manager.get_model_instance("tenant-1", "openai", ModelType.LLM, "gpt-4") + usage = LLMUsage.empty_usage().model_copy(update={"total_tokens": 12}) + result = MagicMock(spec=LLMResult, usage=usage) + reservation = MagicMock(commit_before_delivery=True) + + with ( + patch.object(model_instance, "reserve_quota", return_value=reservation), + patch.object(ModelInstance, "invoke_llm", return_value=result) as invoke, + ): + response = model_instance.invoke_llm(prompt_messages=[], stream=False) + + assert response is result + invoke.assert_called_once() + reservation.commit.assert_called_once_with(usage) + reservation.release.assert_called_once_with() + + +def test_quota_managed_stream_commits_before_first_chunk() -> None: + manager, _ = _build_model_manager_bundle( + provider_type=ProviderType.SYSTEM, + restrict_models=[RestrictModel(model="gpt-4", model_type=ModelType.LLM)], + ) + model_instance = manager.get_model_instance("tenant-1", "openai", ModelType.LLM, "gpt-4") + chunk = LLMResultChunk( + model="gpt-4", + prompt_messages=[], + delta=LLMResultChunkDelta(index=0, message=AssistantPromptMessage(content="hello")), + ) + reservation = MagicMock(commit_before_delivery=True) + events: list[str] = [] + reservation.commit.side_effect = lambda _usage: events.append("commit") + + with ( + patch.object(model_instance, "reserve_quota", return_value=reservation), + patch.object(ModelInstance, "invoke_llm", return_value=(item for item in [chunk])), + ): + response = model_instance.invoke_llm(prompt_messages=[], stream=True) + assert next(response) is chunk + events.append("delivered") + with pytest.raises(StopIteration): + next(response) + + assert events == ["commit", "delivered"] + reservation.release.assert_called_once_with() + + +def test_quota_managed_stream_releases_when_provider_fails_before_first_chunk() -> None: + manager, _ = _build_model_manager_bundle( + provider_type=ProviderType.SYSTEM, + restrict_models=[RestrictModel(model="gpt-4", model_type=ModelType.LLM)], + ) + model_instance = manager.get_model_instance("tenant-1", "openai", ModelType.LLM, "gpt-4") + reservation = MagicMock(commit_before_delivery=True) + + def failing_stream(): + raise RuntimeError("provider failed") + yield + + with ( + patch.object(model_instance, "reserve_quota", return_value=reservation), + patch.object(ModelInstance, "invoke_llm", return_value=failing_stream()), + pytest.raises(RuntimeError, match="provider failed"), + ): + list(model_instance.invoke_llm(prompt_messages=[], stream=True)) + + reservation.commit.assert_not_called() + reservation.release.assert_called_once_with() + + +def test_quota_managed_usage_stream_commits_before_delivering_buffered_chunks() -> None: + manager, _ = _build_model_manager_bundle( + provider_type=ProviderType.SYSTEM, + restrict_models=[RestrictModel(model="gpt-4", model_type=ModelType.LLM)], + ) + model_instance = manager.get_model_instance("tenant-1", "openai", ModelType.LLM, "gpt-4") + usage = LLMUsage.empty_usage().model_copy(update={"total_tokens": 12}) + chunks = [ + LLMResultChunk( + model="gpt-4", + prompt_messages=[], + delta=LLMResultChunkDelta(index=0, message=AssistantPromptMessage(content="hello")), + ), + LLMResultChunk( + model="gpt-4", + prompt_messages=[], + delta=LLMResultChunkDelta(index=1, message=AssistantPromptMessage(content=" world"), usage=usage), + ), + ] + reservation = MagicMock(commit_before_delivery=False) + events: list[str] = [] + reservation.commit.side_effect = lambda _usage: events.append("commit") + + def provider_stream(): + for index, chunk in enumerate(chunks): + events.append(f"provider-{index}") + yield chunk + + with ( + patch.object(model_instance, "reserve_quota", return_value=reservation), + patch.object(ModelInstance, "invoke_llm", return_value=provider_stream()), + ): + response = model_instance.invoke_llm(prompt_messages=[], stream=True) + assert next(response) is chunks[0] + events.append("delivered") + assert list(response) == [chunks[1]] + + assert events == ["provider-0", "provider-1", "commit", "delivered"] + reservation.commit.assert_called_once_with(usage) + reservation.release.assert_called_once_with() + + +def test_quota_managed_usage_stream_does_not_deliver_when_settlement_fails() -> None: + manager, _ = _build_model_manager_bundle( + provider_type=ProviderType.SYSTEM, + restrict_models=[RestrictModel(model="gpt-4", model_type=ModelType.LLM)], + ) + model_instance = manager.get_model_instance("tenant-1", "openai", ModelType.LLM, "gpt-4") + chunk = LLMResultChunk( + model="gpt-4", + prompt_messages=[], + delta=LLMResultChunkDelta(index=0, message=AssistantPromptMessage(content="hello")), + ) + reservation = MagicMock(commit_before_delivery=False) + reservation.commit.side_effect = ValueError("terminal usage is required") + + with ( + patch.object(model_instance, "reserve_quota", return_value=reservation), + patch.object(ModelInstance, "invoke_llm", return_value=(item for item in [chunk])), + pytest.raises(ValueError, match="terminal usage is required"), + ): + next(model_instance.invoke_llm(prompt_messages=[], stream=True)) + + reservation.release.assert_called_once_with() + + def test_lb_model_manager_fetch_next(mocker: MockerFixture, lb_model_manager: LBModelManager): # initialize redis client redis_client.initialize(redis.Redis()) diff --git a/api/tests/unit_tests/core/workflow/graph_engine/layers/test_llm_quota.py b/api/tests/unit_tests/core/workflow/graph_engine/layers/test_llm_quota.py deleted file mode 100644 index 97d7e4a937b..00000000000 --- a/api/tests/unit_tests/core/workflow/graph_engine/layers/test_llm_quota.py +++ /dev/null @@ -1,353 +0,0 @@ -import logging -import threading -from datetime import datetime -from types import SimpleNamespace -from unittest.mock import MagicMock, patch - -import pytest - -from core.app.workflow.layers.llm_quota import LLMQuotaLayer -from core.errors.error import QuotaExceededError -from graphon.enums import BuiltinNodeTypes, WorkflowNodeExecutionStatus -from graphon.graph_engine.entities.commands import CommandType -from graphon.graph_events import NodeRunSucceededEvent -from graphon.model_runtime.entities.llm_entities import LLMUsage -from graphon.node_events import NodeRunResult - - -def _build_succeeded_event(*, provider: str = "openai", model_name: str = "gpt-4o") -> NodeRunSucceededEvent: - return NodeRunSucceededEvent( - id="execution-id", - node_id="llm-node-id", - node_type=BuiltinNodeTypes.LLM, - start_at=datetime.now(), - node_run_result=NodeRunResult( - status=WorkflowNodeExecutionStatus.SUCCEEDED, - inputs={ - "question": "hello", - "model_provider": provider, - "model_name": model_name, - }, - llm_usage=LLMUsage.empty_usage(), - ), - ) - - -def _build_public_model_identity(*, provider: str = "openai", model_name: str = "gpt-4o") -> SimpleNamespace: - return SimpleNamespace(provider=provider, name=model_name) - - -def _build_node_data(*, model: SimpleNamespace | None = None) -> SimpleNamespace: - return SimpleNamespace( - error_strategy=None, - retry_config=SimpleNamespace(retry_enabled=False), - model=model, - ) - - -def _build_node(*, node_type: BuiltinNodeTypes = BuiltinNodeTypes.LLM) -> MagicMock: - node = MagicMock() - node.id = "node-id" - node.execution_id = "execution-id" - node.node_type = node_type - node.node_data = _build_node_data(model=_build_public_model_identity()) - node.model_instance = SimpleNamespace(provider="stale-provider", model_name="stale-model") - return node - - -class _RunnableQuotaNode: - id = "node-id" - execution_id = "execution-id" - node_type = BuiltinNodeTypes.LLM - title = "LLM node" - - def __init__(self, *, stop_event: threading.Event, node_data: SimpleNamespace | None = None) -> None: - self.node_data = node_data or _build_node_data(model=_build_public_model_identity()) - self.graph_runtime_state = SimpleNamespace(stop_event=stop_event) - self.original_run_called = False - - def _run(self) -> NodeRunResult: - self.original_run_called = True - return NodeRunResult(status=WorkflowNodeExecutionStatus.SUCCEEDED) - - -def test_deduct_quota_called_for_successful_llm_node() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - node = _build_node(node_type=BuiltinNodeTypes.LLM) - result_event = _build_succeeded_event() - - with patch("core.app.workflow.layers.llm_quota.deduct_llm_quota_for_model", autospec=True) as mock_deduct: - layer.on_node_run_end(node=node, error=None, result_event=result_event) - - mock_deduct.assert_called_once_with( - tenant_id="tenant-id", - provider="openai", - model="gpt-4o", - usage=result_event.node_run_result.llm_usage, - ) - - -def test_deduct_quota_called_for_question_classifier_node() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - node = _build_node(node_type=BuiltinNodeTypes.QUESTION_CLASSIFIER) - result_event = _build_succeeded_event(provider="anthropic", model_name="claude-3-7-sonnet") - - with patch("core.app.workflow.layers.llm_quota.deduct_llm_quota_for_model", autospec=True) as mock_deduct: - layer.on_node_run_end(node=node, error=None, result_event=result_event) - - mock_deduct.assert_called_once_with( - tenant_id="tenant-id", - provider="anthropic", - model="claude-3-7-sonnet", - usage=result_event.node_run_result.llm_usage, - ) - - -def test_non_llm_node_is_ignored() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - node = _build_node(node_type=BuiltinNodeTypes.START) - result_event = _build_succeeded_event() - - with patch("core.app.workflow.layers.llm_quota.deduct_llm_quota_for_model", autospec=True) as mock_deduct: - layer.on_node_run_end(node=node, error=None, result_event=result_event) - - mock_deduct.assert_not_called() - - -def test_precheck_ignores_non_quota_node() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - node = _build_node(node_type=BuiltinNodeTypes.START) - - with patch("core.app.workflow.layers.llm_quota.ensure_llm_quota_available_for_model", autospec=True) as mock_check: - layer.on_node_run_start(node) - - mock_check.assert_not_called() - - -def test_quota_error_is_handled_in_layer(caplog: pytest.LogCaptureFixture) -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - stop_event = threading.Event() - layer.command_channel = MagicMock() - - node = _build_node(node_type=BuiltinNodeTypes.LLM) - node.graph_runtime_state = MagicMock() - node.graph_runtime_state.stop_event = stop_event - result_event = _build_succeeded_event() - - with ( - caplog.at_level(logging.ERROR, logger="core.app.workflow.layers.llm_quota"), - patch( - "core.app.workflow.layers.llm_quota.deduct_llm_quota_for_model", - autospec=True, - side_effect=ValueError("quota exceeded"), - ) as mock_deduct, - ): - layer.on_node_run_end(node=node, error=None, result_event=result_event) - - mock_deduct.assert_called_once_with( - tenant_id="tenant-id", - provider="openai", - model="gpt-4o", - usage=result_event.node_run_result.llm_usage, - ) - assert "LLM quota deduction failed, node_id=node-id" in caplog.text - assert not stop_event.is_set() - layer.command_channel.send_command.assert_not_called() - - -def test_send_abort_command_is_noop_without_channel_or_after_abort() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - - layer._send_abort_command(reason="no channel") - - layer.command_channel = MagicMock() - layer._abort_sent = True - layer._send_abort_command(reason="already aborted") - - layer.command_channel.send_command.assert_not_called() - - -def test_quota_deduction_exceeded_aborts_workflow_immediately() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - stop_event = threading.Event() - layer.command_channel = MagicMock() - - node = _build_node(node_type=BuiltinNodeTypes.LLM) - node.graph_runtime_state = MagicMock() - node.graph_runtime_state.stop_event = stop_event - - result_event = _build_succeeded_event() - with patch( - "core.app.workflow.layers.llm_quota.deduct_llm_quota_for_model", - autospec=True, - side_effect=QuotaExceededError("No credits remaining"), - ): - layer.on_node_run_end(node=node, error=None, result_event=result_event) - - assert stop_event.is_set() - layer.command_channel.send_command.assert_called_once() - abort_command = layer.command_channel.send_command.call_args.args[0] - assert abort_command.command_type == CommandType.ABORT - assert abort_command.reason == "No credits remaining" - - -def test_quota_precheck_failure_aborts_workflow_immediately() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - stop_event = threading.Event() - layer.command_channel = MagicMock() - - node = _build_node(node_type=BuiltinNodeTypes.LLM) - node.graph_runtime_state = MagicMock() - node.graph_runtime_state.stop_event = stop_event - - with patch( - "core.app.workflow.layers.llm_quota.ensure_llm_quota_available_for_model", - autospec=True, - side_effect=QuotaExceededError("Model provider openai quota exceeded."), - ): - layer.on_node_run_start(node) - - assert stop_event.is_set() - layer.command_channel.send_command.assert_called_once() - abort_command = layer.command_channel.send_command.call_args.args[0] - assert abort_command.command_type == CommandType.ABORT - assert abort_command.reason == "Model provider openai quota exceeded." - - -def test_quota_precheck_failure_blocks_current_node_run() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - stop_event = threading.Event() - layer.command_channel = MagicMock() - - node = _RunnableQuotaNode(stop_event=stop_event) - - with patch( - "core.app.workflow.layers.llm_quota.ensure_llm_quota_available_for_model", - autospec=True, - side_effect=QuotaExceededError("Model provider openai quota exceeded."), - ): - layer.on_node_run_start(node) - - result = node._run() - assert not node.original_run_called - assert result.status == WorkflowNodeExecutionStatus.FAILED - assert result.error == "Model provider openai quota exceeded." - assert result.error_type == QuotaExceededError.__name__ - - -def test_missing_model_identity_blocks_current_node_run() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - stop_event = threading.Event() - layer.command_channel = MagicMock() - - node = _RunnableQuotaNode(stop_event=stop_event, node_data=_build_node_data()) - - with patch("core.app.workflow.layers.llm_quota.ensure_llm_quota_available_for_model", autospec=True) as mock_check: - layer.on_node_run_start(node) - - result = node._run() - assert not node.original_run_called - assert result.status == WorkflowNodeExecutionStatus.FAILED - assert result.error == "LLM quota check requires public node model identity before execution." - assert result.error_type == "LLMQuotaIdentityError" - mock_check.assert_not_called() - - -def test_quota_precheck_passes_without_abort() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - stop_event = threading.Event() - layer.command_channel = MagicMock() - - node = _build_node(node_type=BuiltinNodeTypes.LLM) - node.graph_runtime_state = MagicMock() - node.graph_runtime_state.stop_event = stop_event - - with patch("core.app.workflow.layers.llm_quota.ensure_llm_quota_available_for_model", autospec=True) as mock_check: - layer.on_node_run_start(node) - - assert not stop_event.is_set() - mock_check.assert_called_once_with( - tenant_id="tenant-id", - provider="openai", - model="gpt-4o", - ) - layer.command_channel.send_command.assert_not_called() - - -def test_precheck_reads_model_identity_from_data_when_node_data_is_absent() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - node = SimpleNamespace( - id="node-id", - node_type=BuiltinNodeTypes.LLM, - data=_build_node_data(model=_build_public_model_identity(provider="anthropic", model_name="claude")), - ) - - with patch("core.app.workflow.layers.llm_quota.ensure_llm_quota_available_for_model", autospec=True) as mock_check: - layer.on_node_run_start(node) - - mock_check.assert_called_once_with( - tenant_id="tenant-id", - provider="anthropic", - model="claude", - ) - - -def test_precheck_rejects_invalid_public_model_identity() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - stop_event = threading.Event() - layer.command_channel = MagicMock() - - node = _build_node(node_type=BuiltinNodeTypes.LLM) - node.node_data = _build_node_data(model=_build_public_model_identity(provider="", model_name="gpt-4o")) - node.graph_runtime_state = MagicMock() - node.graph_runtime_state.stop_event = stop_event - - with patch("core.app.workflow.layers.llm_quota.ensure_llm_quota_available_for_model", autospec=True) as mock_check: - layer.on_node_run_start(node) - - assert stop_event.is_set() - mock_check.assert_not_called() - layer.command_channel.send_command.assert_called_once() - - -def test_precheck_requires_public_node_model_config() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - stop_event = threading.Event() - layer.command_channel = MagicMock() - - node = _build_node(node_type=BuiltinNodeTypes.LLM) - node.node_data = _build_node_data() - node.graph_runtime_state = MagicMock() - node.graph_runtime_state.stop_event = stop_event - - with patch("core.app.workflow.layers.llm_quota.ensure_llm_quota_available_for_model", autospec=True) as mock_check: - layer.on_node_run_start(node) - - assert stop_event.is_set() - mock_check.assert_not_called() - layer.command_channel.send_command.assert_called_once() - abort_command = layer.command_channel.send_command.call_args.args[0] - assert abort_command.command_type == CommandType.ABORT - assert abort_command.reason == "LLM quota check requires public node model identity before execution." - - -def test_deduction_requires_public_event_model_identity() -> None: - layer = LLMQuotaLayer(tenant_id="tenant-id") - stop_event = threading.Event() - layer.command_channel = MagicMock() - - node = _build_node(node_type=BuiltinNodeTypes.LLM) - node.graph_runtime_state = MagicMock() - node.graph_runtime_state.stop_event = stop_event - result_event = _build_succeeded_event() - result_event.node_run_result.inputs = {"question": "hello"} - - with patch("core.app.workflow.layers.llm_quota.deduct_llm_quota_for_model", autospec=True) as mock_deduct: - layer.on_node_run_end(node=node, error=None, result_event=result_event) - - assert stop_event.is_set() - mock_deduct.assert_not_called() - layer.command_channel.send_command.assert_called_once() - abort_command = layer.command_channel.send_command.call_args.args[0] - assert abort_command.command_type == CommandType.ABORT - assert abort_command.reason == "LLM quota deduction requires model identity in the node result event." diff --git a/api/tests/unit_tests/core/workflow/test_llm_node.py b/api/tests/unit_tests/core/workflow/test_llm_node.py new file mode 100644 index 00000000000..85a77dfd59e --- /dev/null +++ b/api/tests/unit_tests/core/workflow/test_llm_node.py @@ -0,0 +1,56 @@ +from collections.abc import Generator +from types import SimpleNamespace +from typing import cast +from unittest.mock import Mock, sentinel + +import pytest + +from core.workflow.llm_node import DifyLLMNode +from graphon.nodes.llm.node import LLMNode +from graphon.nodes.llm.runtime_protocols import LLMPollingCapableProtocol + + +def test_dify_llm_node_finalizes_polling_when_generator_is_closed(monkeypatch: pytest.MonkeyPatch) -> None: + def invoke(*args: object, **kwargs: object) -> Generator[object, None, None]: + _ = args, kwargs + yield sentinel.event + yield sentinel.unconsumed + + monkeypatch.setattr(LLMNode, "_invoke_llm_with_polling", invoke) + finalizer = Mock() + node = object.__new__(DifyLLMNode) + node._polling_finalizer = finalizer + + events = node._invoke_llm_with_polling( + polling_model=cast(LLMPollingCapableProtocol, SimpleNamespace()), + prompt_messages=[], + stop=None, + ) + + assert next(events) is sentinel.event + events.close() + + finalizer.assert_called_once_with() + + +def test_dify_llm_node_finalizes_polling_when_polling_fails(monkeypatch: pytest.MonkeyPatch) -> None: + def invoke(*args: object, **kwargs: object) -> Generator[object, None, None]: + _ = args, kwargs + yield sentinel.event + raise RuntimeError("polling failed") + + monkeypatch.setattr(LLMNode, "_invoke_llm_with_polling", invoke) + finalizer = Mock() + node = object.__new__(DifyLLMNode) + node._polling_finalizer = finalizer + events = node._invoke_llm_with_polling( + polling_model=cast(LLMPollingCapableProtocol, SimpleNamespace()), + prompt_messages=[], + stop=None, + ) + + assert next(events) is sentinel.event + with pytest.raises(RuntimeError, match="polling failed"): + next(events) + + finalizer.assert_called_once_with() diff --git a/api/tests/unit_tests/core/workflow/test_node_factory.py b/api/tests/unit_tests/core/workflow/test_node_factory.py index 0e055cc2c40..f5dd034a649 100644 --- a/api/tests/unit_tests/core/workflow/test_node_factory.py +++ b/api/tests/unit_tests/core/workflow/test_node_factory.py @@ -12,6 +12,7 @@ from core.plugin.impl.model_runtime import PluginModelRuntime from core.plugin.plugin_service import PluginService from core.workflow import node_factory from core.workflow import template_rendering as workflow_template_rendering +from core.workflow.llm_node import DifyLLMNode from core.workflow.node_runtime import DifyPreparedLLM from core.workflow.nodes.knowledge_index import KNOWLEDGE_INDEX_NODE_TYPE from graphon.entities.base_node_data import BaseNodeData @@ -688,7 +689,7 @@ class TestDifyNodeFactoryCreateNode: }, } ) - wrapped_model_instance = sentinel.wrapped_model_instance + wrapped_model_instance = MagicMock(spec=DifyPreparedLLM) memory = sentinel.memory factory._build_model_instance_for_llm_node = MagicMock(return_value=sentinel.model_instance) factory._build_memory_for_llm_node = MagicMock(return_value=memory) @@ -717,6 +718,7 @@ class TestDifyNodeFactoryCreateNode: request_metadata={"app_id": "app-id"}, ) assert kwargs["model_instance"] is wrapped_model_instance + assert kwargs["polling_finalizer"] is wrapped_model_instance.finalize_llm_polling def test_resolve_llm_model_reference_uses_shared_model_and_parameters(self, factory): node_data = LLMNodeData.model_validate( @@ -971,6 +973,44 @@ class TestDifyNodeFactoryCreateNode: assert node.node_data.structured_output_switch_on is True assert node.node_data.structured_output_enabled is True + def test_create_node_uses_dify_llm_node_for_persisted_version_one(self, monkeypatch, factory): + factory.graph_init_params = SimpleNamespace( + workflow_id="workflow-id", + graph_config={}, + run_context={}, + call_depth=0, + ) + monkeypatch.setattr( + factory, + "_build_llm_compatible_node_init_kwargs", + MagicMock( + return_value={ + "model_instance": sentinel.model_instance, + "llm_file_saver": sentinel.llm_file_saver, + "prompt_message_serializer": sentinel.prompt_message_serializer, + "polling_finalizer": MagicMock(), + } + ), + ) + + node = factory.create_node( + { + "id": "llm-node-id", + "data": { + "type": BuiltinNodeTypes.LLM, + "version": "1", + "title": "LLM", + "model": {"provider": "provider", "name": "model", "mode": "chat"}, + "prompt_template": [{"role": "system", "text": "x"}], + "context": {"enabled": False, "variable_selector": []}, + "vision": {"enabled": False}, + }, + } + ) + + assert isinstance(node, DifyLLMNode) + assert node.version() == "1" + @pytest.mark.parametrize( ("node_type", "constructor_name", "expected_extra_kwargs"), [ diff --git a/api/tests/unit_tests/core/workflow/test_node_runtime.py b/api/tests/unit_tests/core/workflow/test_node_runtime.py index adfd7ed2c5f..39cfa96c243 100644 --- a/api/tests/unit_tests/core/workflow/test_node_runtime.py +++ b/api/tests/unit_tests/core/workflow/test_node_runtime.py @@ -11,6 +11,7 @@ from sqlalchemy.orm import Session, sessionmaker from core.app.entities.app_invoke_entities import DIFY_RUN_CONTEXT_KEY, DifyRunContext, InvokeFrom, UserFrom from core.app.file_access import FileAccessScope, bind_file_access_scope, grant_retriever_segment_access from core.llm_generator.output_parser.errors import OutputParserError +from core.model_manager import QuotaManagedModelInstance from core.plugin.impl.exc import PluginLLMPollingUnsupportedError from core.plugin.impl.model import PluginModelClient from core.plugin.impl.model_runtime import PluginModelRuntime @@ -148,6 +149,12 @@ class _ModelInstanceStub: self.invoke_llm = Mock(return_value=invoke_llm_result) +class _QuotaManagedModelInstanceStub(_ModelInstanceStub, QuotaManagedModelInstance): + def __init__(self, **kwargs) -> None: + super().__init__(**kwargs) + self.reserve_quota = Mock() + + def _build_run_context(*, invoke_from: InvokeFrom | str = InvokeFrom.DEBUGGER) -> dict[str, object]: return build_test_run_context( tenant_id="tenant-id", @@ -357,6 +364,146 @@ def test_dify_prepared_polling_llm_delegates_to_plugin_runtime() -> None: ) +def test_dify_prepared_polling_llm_commits_successful_reservation() -> None: + running_result = LLMPollingResult( + status=LLMPollingStatus.RUNNING, + plugin_state={"task_id": "poll-1"}, + ) + usage = node_runtime.LLMUsage.empty_usage().model_copy(update={"total_tokens": 5}) + succeeded_result = LLMPollingResult( + status=LLMPollingStatus.SUCCEEDED, + result=node_runtime.LLMResult( + model="gpt-4o-mini", + prompt_messages=[], + message=AssistantPromptMessage(content="done"), + usage=usage, + ), + ) + plugin_runtime = PluginModelRuntime( + tenant_id="tenant-id", + user_id="user-id", + client=Mock(spec=PluginModelClient), + plugin_service=PluginService, + ) + plugin_runtime.start_llm_polling = Mock(return_value=running_result) # type: ignore[method-assign] + plugin_runtime.check_llm_polling = Mock(return_value=succeeded_result) # type: ignore[method-assign] + reservation = MagicMock() + model_instance = _QuotaManagedModelInstanceStub( + model_schema=_build_model_schema(features=[ModelFeature.POLLING]), + model_runtime=plugin_runtime, + ) + model_instance.reserve_quota.return_value = reservation + prepared = DifyPreparedPollingLLM(model_instance) + + prepared.start_llm_polling( + prompt_messages=[], + model_parameters={}, + tools=None, + stop=None, + json_schema=None, + ) + prepared.check_llm_polling(plugin_state={"task_id": "poll-1"}) + + reservation.commit.assert_called_once_with(usage) + reservation.release.assert_not_called() + + +def test_dify_prepared_polling_llm_releases_previous_reservation_on_restart() -> None: + running_result = LLMPollingResult( + status=LLMPollingStatus.RUNNING, + plugin_state={"task_id": "poll-1"}, + ) + plugin_runtime = PluginModelRuntime( + tenant_id="tenant-id", + user_id="user-id", + client=Mock(spec=PluginModelClient), + plugin_service=PluginService, + ) + plugin_runtime.start_llm_polling = Mock(return_value=running_result) # type: ignore[method-assign] + first_reservation = MagicMock() + second_reservation = MagicMock() + model_instance = _QuotaManagedModelInstanceStub( + model_schema=_build_model_schema(features=[ModelFeature.POLLING]), + model_runtime=plugin_runtime, + ) + model_instance.reserve_quota.side_effect = [first_reservation, second_reservation] + prepared = DifyPreparedPollingLLM(model_instance) + + for _ in range(2): + prepared.start_llm_polling( + prompt_messages=[], + model_parameters={}, + tools=None, + stop=None, + json_schema=None, + ) + + first_reservation.release.assert_called_once_with() + second_reservation.release.assert_not_called() + assert model_instance.reserve_quota.call_count == 2 + + +def test_dify_prepared_polling_llm_releases_reservation_when_finalized() -> None: + running_result = LLMPollingResult( + status=LLMPollingStatus.RUNNING, + plugin_state={"task_id": "poll-1"}, + ) + polling_runtime = SimpleNamespace( + start_llm_polling=Mock(return_value=running_result), + check_llm_polling=Mock(), + ) + reservation = MagicMock() + model_instance = _QuotaManagedModelInstanceStub( + model_schema=_build_model_schema(features=[ModelFeature.POLLING]), + model_runtime=polling_runtime, + ) + model_instance.reserve_quota.return_value = reservation + prepared = DifyPreparedPollingLLM(model_instance) + + prepared.start_llm_polling( + prompt_messages=[], + model_parameters={}, + tools=None, + stop=None, + json_schema=None, + ) + prepared.finalize_llm_polling() + prepared.finalize_llm_polling() + + reservation.release.assert_called_once_with() + + +def test_dify_prepared_polling_llm_releases_reservation_when_check_fails() -> None: + running_result = LLMPollingResult( + status=LLMPollingStatus.RUNNING, + plugin_state={"task_id": "poll-1"}, + ) + polling_runtime = SimpleNamespace( + start_llm_polling=Mock(return_value=running_result), + check_llm_polling=Mock(side_effect=RuntimeError("polling failed")), + ) + reservation = MagicMock() + model_instance = _QuotaManagedModelInstanceStub( + model_schema=_build_model_schema(features=[ModelFeature.POLLING]), + model_runtime=polling_runtime, + ) + model_instance.reserve_quota.return_value = reservation + prepared = DifyPreparedPollingLLM(model_instance) + + prepared.start_llm_polling( + prompt_messages=[], + model_parameters={}, + tools=None, + stop=None, + json_schema=None, + ) + + with pytest.raises(RuntimeError, match="polling failed"): + prepared.check_llm_polling(plugin_state={"task_id": "poll-1"}) + + reservation.release.assert_called_once_with() + + def test_dify_prepared_polling_llm_raise_exception_when_polling_is_unsupported() -> None: llm_result = node_runtime.LLMResult( model="gpt-4o-mini", diff --git a/api/tests/unit_tests/core/workflow/test_workflow_entry_helpers.py b/api/tests/unit_tests/core/workflow/test_workflow_entry_helpers.py index 6614d9386cf..847acee37e5 100644 --- a/api/tests/unit_tests/core/workflow/test_workflow_entry_helpers.py +++ b/api/tests/unit_tests/core/workflow/test_workflow_entry_helpers.py @@ -1,4 +1,3 @@ -import threading from collections import UserString from datetime import datetime from types import SimpleNamespace @@ -16,7 +15,6 @@ from graphon.errors import WorkflowNodeRunFailedError from graphon.file import File, FileTransferMethod, FileType from graphon.filters import ResponseStreamFilter from graphon.graph_events import GraphRunFailedEvent, NodeRunSucceededEvent -from graphon.model_runtime.entities.llm_entities import LLMUsage from graphon.node_events import NodeRunResult from graphon.nodes import BuiltinNodeTypes from graphon.runtime import VariablePool @@ -40,7 +38,6 @@ def _build_minimal_workflow_entry( monkeypatch.setattr(workflow_entry, "GraphEngine", MagicMock(return_value=graph_engine)) monkeypatch.setattr(workflow_entry, "GraphEngineConfig", MagicMock(return_value=sentinel.graph_engine_config)) monkeypatch.setattr(workflow_entry, "InMemoryChannel", MagicMock(return_value=sentinel.command_channel)) - monkeypatch.setattr(workflow_entry, "LLMQuotaLayer", MagicMock(return_value=sentinel.llm_quota_layer)) return workflow_entry.WorkflowEntry( tenant_id="tenant-id", @@ -82,7 +79,6 @@ class TestWorkflowEntryInit: graph_runtime_state = SimpleNamespace(_execution_context=None) debug_layer = sentinel.debug_layer execution_limits_layer = sentinel.execution_limits_layer - llm_quota_layer = sentinel.llm_quota_layer observability_layer = sentinel.observability_layer with ( @@ -99,7 +95,6 @@ class TestWorkflowEntryInit: "ExecutionLimitsLayer", return_value=execution_limits_layer, ) as execution_limits_layer_cls, - patch.object(workflow_entry, "LLMQuotaLayer", return_value=llm_quota_layer) as llm_quota_layer_cls, patch.object(workflow_entry, "ObservabilityLayer", return_value=observability_layer), ): entry = workflow_entry.WorkflowEntry( @@ -137,11 +132,9 @@ class TestWorkflowEntryInit: max_steps=workflow_entry.dify_config.WORKFLOW_MAX_EXECUTION_STEPS, max_time=workflow_entry.dify_config.WORKFLOW_MAX_EXECUTION_TIME, ) - llm_quota_layer_cls.assert_called_once_with(tenant_id="tenant-id") assert graph_engine.layer.call_args_list == [ ((debug_layer,), {}), ((execution_limits_layer,), {}), - ((llm_quota_layer,), {}), ((observability_layer,), {}), ] @@ -754,7 +747,6 @@ class TestMappingUserInputsBranches: class TestWorkflowEntryNodeLayers: def test_run_node_with_layers_reports_success(self): - quota_layer = MagicMock() observability_layer = MagicMock() result_event = NodeRunSucceededEvent( id="execution-id", @@ -775,7 +767,6 @@ class TestWorkflowEntryNodeLayers: node = FakeNode() with ( - patch.object(workflow_entry, "LLMQuotaLayer", return_value=quota_layer) as quota_layer_cls, patch.object(workflow_entry, "ObservabilityLayer", return_value=observability_layer), patch.object(workflow_entry, "InMemoryChannel", return_value=sentinel.command_channel), patch.object( @@ -787,9 +778,8 @@ class TestWorkflowEntryNodeLayers: events = list(workflow_entry.WorkflowEntry._run_node_with_layers(node, tenant_id="tenant-id")) assert events == [result_event] - quota_layer_cls.assert_called_once_with(tenant_id="tenant-id") runtime_state_wrapper.assert_called_once_with(sentinel.graph_runtime_state) - for layer in (quota_layer, observability_layer): + for layer in (observability_layer,): layer.initialize.assert_called_once_with(sentinel.read_only_runtime_state, sentinel.command_channel) layer.on_graph_start.assert_called_once_with() layer.on_node_run_start.assert_called_once_with(node) @@ -797,7 +787,6 @@ class TestWorkflowEntryNodeLayers: layer.on_graph_end.assert_called_once_with(None) def test_run_node_with_layers_reports_errors(self): - quota_layer = MagicMock() observability_layer = MagicMock() class FakeNode: @@ -812,7 +801,6 @@ class TestWorkflowEntryNodeLayers: node = FakeNode() with ( - patch.object(workflow_entry, "LLMQuotaLayer", return_value=quota_layer), patch.object(workflow_entry, "ObservabilityLayer", return_value=observability_layer), patch.object( workflow_entry, @@ -823,64 +811,8 @@ class TestWorkflowEntryNodeLayers: with pytest.raises(RuntimeError, match="boom"): list(workflow_entry.WorkflowEntry._run_node_with_layers(node, tenant_id="tenant-id")) - for layer in (quota_layer, observability_layer): + for layer in (observability_layer,): assert layer.on_node_run_end.call_args.args[0] is node assert isinstance(layer.on_node_run_end.call_args.args[1], RuntimeError) assert layer.on_node_run_end.call_args.args[2] is None assert isinstance(layer.on_graph_end.call_args.args[0], RuntimeError) - - def test_run_node_with_layers_deducts_llm_quota(self): - result_event = NodeRunSucceededEvent( - id="execution-id", - node_id="node-id", - node_type=BuiltinNodeTypes.LLM, - start_at=datetime.now(), - node_run_result=NodeRunResult( - status=WorkflowNodeExecutionStatus.SUCCEEDED, - inputs={"model_provider": "openai", "model_name": "gpt-4o"}, - llm_usage=LLMUsage.empty_usage(), - ), - ) - - class FakeNode: - id = "node-id" - node_type = BuiltinNodeTypes.LLM - graph_runtime_state = SimpleNamespace( - stop_event=threading.Event(), - variable_pool=VariablePool(), - ) - node_data = SimpleNamespace( - model=SimpleNamespace(provider="openai", name="gpt-4o"), - error_strategy=None, - retry_config=SimpleNamespace(retry_enabled=False), - ) - - def bind_execution_id(self, execution_id): - self.execution_id = execution_id - - def run(self): - yield result_event - - with ( - patch.object(workflow_entry, "ObservabilityLayer", return_value=MagicMock()), - patch( - "core.app.workflow.layers.llm_quota.ensure_llm_quota_available_for_model", - autospec=True, - ) as ensure_quota, - patch( - "core.app.workflow.layers.llm_quota.deduct_llm_quota_for_model", - autospec=True, - ) as deduct_quota, - ): - generator = workflow_entry.WorkflowEntry._run_node_with_layers(FakeNode(), tenant_id="tenant-id") - event = next(generator) - - assert event is result_event - ensure_quota.assert_called_once_with(tenant_id="tenant-id", provider="openai", model="gpt-4o") - deduct_quota.assert_called_once_with( - tenant_id="tenant-id", - provider="openai", - model="gpt-4o", - usage=result_event.node_run_result.llm_usage, - ) - generator.close() diff --git a/api/tests/unit_tests/services/test_credit_pool_service.py b/api/tests/unit_tests/services/test_credit_pool_service.py index 09008f60141..4e8f64b2ef6 100644 --- a/api/tests/unit_tests/services/test_credit_pool_service.py +++ b/api/tests/unit_tests/services/test_credit_pool_service.py @@ -17,6 +17,7 @@ from services.credit_pool_service import ( CREDIT_POOL_TENANT_LOCK_TIMEOUT_SECONDS, FEATURE_KEY_CREDIT_POOL, CreditPoolBalance, + CreditPoolReservationState, CreditPoolService, ) @@ -272,6 +273,90 @@ def test_get_pool_uses_billing_quota_balance_when_enabled() -> None: ) +def test_reserve_credits_commits_billing_reservation_once() -> None: + with ( + patch.object(CreditPoolService, "_use_billing_quota", return_value=True), + patch("services.billing_service.BillingService.quota_reserve") as quota_reserve, + patch("services.billing_service.BillingService.quota_commit") as quota_commit, + patch("services.billing_service.BillingService.quota_release") as quota_release, + ): + quota_reserve.return_value = {"reservation_id": "reservation-1", "available": 7, "reserved": 3} + + reservation = CreditPoolService.reserve_credits( + tenant_id="tenant-1", + credits_required=3, + pool_type=ProviderQuotaType.TRIAL, + request_id="request-1", + meta={"source": "test"}, + ) + reservation.commit() + reservation.commit() + reservation.release() + + assert reservation.state == CreditPoolReservationState.COMMITTED + quota_reserve.assert_called_once_with( + tenant_id="tenant-1", + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket="trial", + request_id="request-1", + amount=3, + meta={"source": "test"}, + ) + quota_commit.assert_called_once_with( + tenant_id="tenant-1", + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket="trial", + reservation_id="reservation-1", + actual_amount=3, + meta={"source": "test", "request_id": "request-1"}, + ) + quota_release.assert_not_called() + + +def test_reserve_credits_releases_billing_reservation() -> None: + with ( + patch.object(CreditPoolService, "_use_billing_quota", return_value=True), + patch("services.billing_service.BillingService.quota_reserve") as quota_reserve, + patch("services.billing_service.BillingService.quota_release") as quota_release, + ): + quota_reserve.return_value = {"reservation_id": "reservation-1", "available": 7, "reserved": 3} + + reservation = CreditPoolService.reserve_credits( + tenant_id="tenant-1", + credits_required=3, + request_id="request-1", + ) + reservation.release() + reservation.release() + + assert reservation.state == CreditPoolReservationState.RELEASED + quota_release.assert_called_once_with( + tenant_id="tenant-1", + feature_key=FEATURE_KEY_CREDIT_POOL, + bucket="trial", + reservation_id="reservation-1", + ) + + +def test_reserve_credits_database_fallback_restores_released_amount(sqlite_session: Session) -> None: + pool = _create_pool(sqlite_session, quota_limit=10, quota_used=2) + redis_lock = _make_redis_lock() + + with patch("services.credit_pool_service.redis_client.lock", return_value=redis_lock): + reservation = CreditPoolService.reserve_credits( + tenant_id=pool.tenant_id, + credits_required=3, + request_id="request-1", + session_factory=lambda: sqlite_session, + ) + assert _get_quota_used(session=sqlite_session, pool_id=pool.id) == 5 + + reservation.release() + + assert reservation.state == CreditPoolReservationState.RELEASED + assert _get_quota_used(session=sqlite_session, pool_id=pool.id) == 2 + + def test_check_and_deduct_credits_uses_billing_reserve_and_commit_when_enabled() -> None: tenant_id = "tenant-1" with (