feat(api): track credit usage contexts (#41181)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
林玮 (Jade Lin)
2026-08-25 10:21:32 +00:00
committed by GitHub
parent 8abff1b0b1
commit 960018253b
72 changed files with 1206 additions and 165 deletions
+21 -12
View File
@@ -12,6 +12,8 @@ from core.agent.errors import AgentMaxIterationError
from core.agent.output_parser.cot_output_parser import CotAgentOutputParser
from core.app.apps.base_app_queue_manager import PublishFrom
from core.app.entities.queue_entities import QueueAgentThoughtEvent, QueueMessageEndEvent, QueueMessageFileEvent
from core.credit_usage import CreditUsageAppType, CreditUsageCreatedBy
from core.model_context import use_credit_usage_metadata
from core.ops.ops_trace_manager import TraceQueueManager
from core.prompt.agent_history_prompt_transform import AgentHistoryPromptTransform
from core.tools.__base.tool import Tool
@@ -135,6 +137,12 @@ class CotAgentRunner(BaseAgentRunner, ABC):
session.close()
# invoke model
request_metadata: dict[str, object] = {
"app_id": self.app_config.app_id,
"app_type": CreditUsageAppType.AGENT,
"created_by": CreditUsageCreatedBy.APP,
}
chunks = model_instance.invoke_llm(
prompt_messages=prompt_messages,
model_parameters=app_generate_entity.model_conf.parameters,
@@ -142,7 +150,7 @@ class CotAgentRunner(BaseAgentRunner, ABC):
stop=app_generate_entity.model_conf.stop,
stream=True,
callbacks=[],
request_metadata={"app_id": self.app_config.app_id},
request_metadata=request_metadata,
)
usage_dict: dict[str, LLMUsage | None] = {}
@@ -331,17 +339,18 @@ class CotAgentRunner(BaseAgentRunner, ABC):
pass
# invoke tool
tool_invoke_response, message_files, tool_invoke_meta = ToolEngine.agent_invoke(
session=session,
tool=tool_instance,
tool_parameters=tool_call_args,
user_id=self.user_id,
tenant_id=self.tenant_id,
message=self.message,
invoke_from=self.application_generate_entity.invoke_from,
agent_tool_callback=self.agent_callback,
trace_manager=trace_manager,
)
with use_credit_usage_metadata({"app_type": CreditUsageAppType.AGENT}):
tool_invoke_response, message_files, tool_invoke_meta = ToolEngine.agent_invoke(
session=session,
tool=tool_instance,
tool_parameters=tool_call_args,
user_id=self.user_id,
tenant_id=self.tenant_id,
message=self.message,
invoke_from=self.application_generate_entity.invoke_from,
agent_tool_callback=self.agent_callback,
trace_manager=trace_manager,
)
session.commit()
session.close()
+24 -15
View File
@@ -13,6 +13,8 @@ from core.agent.errors import AgentMaxIterationError
from core.app.apps.base_app_queue_manager import PublishFrom
from core.app.entities.queue_entities import QueueAgentThoughtEvent, QueueMessageEndEvent, QueueMessageFileEvent
from core.app.file_access import grant_upload_file_access
from core.credit_usage import CreditUsageAppType, CreditUsageCreatedBy
from core.model_context import use_credit_usage_metadata
from core.prompt.agent_history_prompt_transform import AgentHistoryPromptTransform
from core.tools.entities.tool_entities import ToolInvokeMeta
from core.tools.signature import sign_upload_file_preview_url
@@ -167,6 +169,12 @@ class FunctionCallAgentRunner(BaseAgentRunner):
session.close()
# invoke model
request_metadata: dict[str, object] = {
"app_id": self.app_config.app_id,
"app_type": CreditUsageAppType.AGENT,
"created_by": CreditUsageCreatedBy.APP,
}
chunks: Union[Generator[LLMResultChunk, None, None], LLMResult] = model_instance.invoke_llm(
prompt_messages=prompt_messages,
model_parameters=app_generate_entity.model_conf.parameters,
@@ -174,7 +182,7 @@ class FunctionCallAgentRunner(BaseAgentRunner):
stop=app_generate_entity.model_conf.stop,
stream=self.stream_tool_call,
callbacks=[],
request_metadata={"app_id": self.app_config.app_id},
request_metadata=request_metadata,
)
tool_calls: list[tuple[str, str, dict[str, Any]]] = []
@@ -315,20 +323,21 @@ class FunctionCallAgentRunner(BaseAgentRunner):
}
else:
# invoke tool
tool_invoke_response, message_files, tool_invoke_meta = ToolEngine.agent_invoke(
session=session,
tool=tool_instance,
tool_parameters=tool_call_args,
user_id=self.user_id,
tenant_id=self.tenant_id,
message=self.message,
invoke_from=self.application_generate_entity.invoke_from,
agent_tool_callback=self.agent_callback,
trace_manager=trace_manager,
app_id=self.application_generate_entity.app_config.app_id,
message_id=self.message.id,
conversation_id=self.conversation.id,
)
with use_credit_usage_metadata({"app_type": CreditUsageAppType.AGENT}):
tool_invoke_response, message_files, tool_invoke_meta = ToolEngine.agent_invoke(
session=session,
tool=tool_instance,
tool_parameters=tool_call_args,
user_id=self.user_id,
tenant_id=self.tenant_id,
message=self.message,
invoke_from=self.application_generate_entity.invoke_from,
agent_tool_callback=self.agent_callback,
trace_manager=trace_manager,
app_id=self.application_generate_entity.app_config.app_id,
message_id=self.message.id,
conversation_id=self.conversation.id,
)
session.commit()
session.close()
# publish files
@@ -20,6 +20,7 @@ from core.app.entities.app_invoke_entities import (
AppGenerateEntity,
DifyRunContext,
InvokeFrom,
get_credit_usage_app_type,
)
from core.app.entities.queue_entities import (
QueueAnnotationReplyEvent,
@@ -141,6 +142,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
user_id=self.application_generate_entity.user_id,
invoke_from=invoke_from,
user_from=user_from,
app_type=get_credit_usage_app_type(app_config.app_mode),
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
)
elif self.application_generate_entity.single_iteration_run or self.application_generate_entity.single_loop_run:
@@ -150,6 +152,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
single_iteration_run=self.application_generate_entity.single_iteration_run,
single_loop_run=self.application_generate_entity.single_loop_run,
user_id=self.application_generate_entity.user_id,
app_type=get_credit_usage_app_type(app_config.app_mode),
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
)
else:
@@ -220,6 +223,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
user_from=user_from,
invoke_from=invoke_from,
root_node_id=root_node_id,
app_type=get_credit_usage_app_type(app_config.app_mode),
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
)
@@ -280,6 +284,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
user_id=self.application_generate_entity.user_id,
user_from=user_from,
invoke_from=invoke_from,
app_type=get_credit_usage_app_type(app_config.app_mode),
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
)
)
@@ -18,6 +18,7 @@ from core.app.apps.draft_variable_saver import DraftVariableSaverFactory
from core.app.entities.app_invoke_entities import (
AdvancedChatAppGenerateEntity,
InvokeFrom,
get_credit_usage_app_type,
)
from core.app.entities.queue_entities import (
MessageQueueMessage,
@@ -371,7 +372,10 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
and features_dict["text_to_speech"].get("autoPlay") == "enabled"
):
tts_publisher = AppGeneratorTTSPublisher(
tenant_id, features_dict["text_to_speech"].get("voice"), features_dict["text_to_speech"].get("language")
tenant_id,
features_dict["text_to_speech"].get("voice"),
features_dict["text_to_speech"].get("language"),
get_credit_usage_app_type(self._application_generate_entity.app_config.app_mode),
)
try:
@@ -45,6 +45,7 @@ from core.app.entities.app_invoke_entities import (
InvokeFrom,
UserFrom,
)
from core.credit_usage import CreditUsageAppType
from core.db.session_factory import session_factory
from core.ops.ops_trace_manager import TraceQueueManager
from core.workflow.file_reference import build_file_reference, is_canonical_file_reference
@@ -498,6 +499,7 @@ class AgentAppGenerator(MessageBasedAppGenerator):
user_id=application_generate_entity.user_id,
user_from=user_from,
invoke_from=application_generate_entity.invoke_from,
app_type=CreditUsageAppType.AGENT_V2,
)
with session_factory.create_session() as session:
agent, config_version, agent_soul = self._resolve_agent_by_id(
+7 -1
View File
@@ -9,6 +9,8 @@ from core.app.apps.base_app_runner import AppRunner
from core.app.apps.chat.app_config_manager import ChatAppConfig
from core.app.entities.app_invoke_entities import (
ChatAppGenerateEntity,
get_credit_usage_app_type,
get_credit_usage_created_by,
)
from core.app.entities.queue_entities import QueueAnnotationReplyEvent
from core.callback_handler.index_tool_callback_handler import DatasetIndexToolCallbackHandler
@@ -231,12 +233,16 @@ class ChatAppRunner(AppRunner):
model=application_generate_entity.model_conf.model,
)
request_metadata: dict[str, object] = {"app_id": app_config.app_id}
request_metadata["app_type"] = get_credit_usage_app_type(app_config.app_mode)
request_metadata["created_by"] = get_credit_usage_created_by(app_config.app_mode)
invoke_result = model_instance.invoke_llm(
prompt_messages=prompt_messages,
model_parameters=application_generate_entity.model_conf.parameters,
stop=stop,
stream=application_generate_entity.stream,
request_metadata={"app_id": app_config.app_id},
request_metadata=request_metadata,
)
# handle invoke result
+7 -1
View File
@@ -9,6 +9,8 @@ from core.app.apps.base_app_runner import AppRunner
from core.app.apps.completion.app_config_manager import CompletionAppConfig
from core.app.entities.app_invoke_entities import (
CompletionAppGenerateEntity,
get_credit_usage_app_type,
get_credit_usage_created_by,
)
from core.callback_handler.index_tool_callback_handler import DatasetIndexToolCallbackHandler
from core.db.session_factory import create_session
@@ -192,12 +194,16 @@ class CompletionAppRunner(AppRunner):
model=application_generate_entity.model_conf.model,
)
request_metadata: dict[str, object] = {"app_id": app_config.app_id}
request_metadata["app_type"] = get_credit_usage_app_type(app_config.app_mode)
request_metadata["created_by"] = get_credit_usage_created_by(app_config.app_mode)
invoke_result = model_instance.invoke_llm(
prompt_messages=prompt_messages,
model_parameters=application_generate_entity.model_conf.parameters,
stop=stop,
stream=application_generate_entity.stream,
request_metadata={"app_id": app_config.app_id},
request_metadata=request_metadata,
)
# handle invoke result
@@ -15,6 +15,7 @@ from core.app.entities.app_invoke_entities import (
build_dify_run_context,
)
from core.app.workflow.layers.persistence import PersistenceWorkflowInfo, WorkflowPersistenceLayer
from core.credit_usage import CreditUsageAppType
from core.db.session_factory import create_session
from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository
from core.workflow.node_factory import DifyGraphInitContext, DifyNodeFactory, get_default_root_node_id
@@ -302,6 +303,7 @@ class PipelineRunner(WorkflowBasedAppRunner):
user_id=self.application_generate_entity.user_id,
user_from=user_from,
invoke_from=invoke_from,
app_type=CreditUsageAppType.RAG_PIPELINE,
)
graph_init_context = DifyGraphInitContext(
workflow_id=workflow.id,
+10 -1
View File
@@ -13,7 +13,12 @@ from core.app.apps.workflow.command_channels import (
)
from core.app.apps.workflow.stop_aware_ready_queue import attach_stop_aware_ready_queue
from core.app.apps.workflow_app_runner import WorkflowBasedAppRunner
from core.app.entities.app_invoke_entities import DifyRunContext, InvokeFrom, WorkflowAppGenerateEntity
from core.app.entities.app_invoke_entities import (
DifyRunContext,
InvokeFrom,
WorkflowAppGenerateEntity,
get_credit_usage_app_type,
)
from core.app.workflow.layers.persistence import PersistenceWorkflowInfo, WorkflowPersistenceLayer
from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository
from core.workflow.node_factory import get_default_root_node_id
@@ -99,6 +104,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
user_from=user_from,
invoke_from=invoke_from,
root_node_id=self._root_node_id,
app_type=get_credit_usage_app_type(app_config.app_mode),
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
)
elif self.application_generate_entity.single_iteration_run or self.application_generate_entity.single_loop_run:
@@ -107,6 +113,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
single_iteration_run=self.application_generate_entity.single_iteration_run,
single_loop_run=self.application_generate_entity.single_loop_run,
user_id=self.application_generate_entity.user_id,
app_type=get_credit_usage_app_type(app_config.app_mode),
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
)
else:
@@ -150,6 +157,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
user_from=user_from,
invoke_from=invoke_from,
root_node_id=root_node_id,
app_type=get_credit_usage_app_type(app_config.app_mode),
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
)
@@ -210,6 +218,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
user_id=self.application_generate_entity.user_id,
user_from=user_from,
invoke_from=invoke_from,
app_type=get_credit_usage_app_type(app_config.app_mode),
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
)
)
@@ -10,7 +10,7 @@ from core.app.apps.base_app_queue_manager import AppQueueManager
from core.app.apps.common.graph_runtime_state_support import GraphRuntimeStateSupport
from core.app.apps.common.workflow_response_converter import WorkflowResponseConverter
from core.app.apps.draft_variable_saver import DraftVariableSaverFactory
from core.app.entities.app_invoke_entities import InvokeFrom, WorkflowAppGenerateEntity
from core.app.entities.app_invoke_entities import InvokeFrom, WorkflowAppGenerateEntity, get_credit_usage_app_type
from core.app.entities.queue_entities import (
AppQueueEvent,
MessageQueueMessage,
@@ -268,7 +268,10 @@ class WorkflowAppGenerateTaskPipeline(GraphRuntimeStateSupport):
and features_dict["text_to_speech"].get("autoPlay") == "enabled"
):
tts_publisher = AppGeneratorTTSPublisher(
tenant_id, features_dict["text_to_speech"].get("voice"), features_dict["text_to_speech"].get("language")
tenant_id,
features_dict["text_to_speech"].get("voice"),
features_dict["text_to_speech"].get("language"),
get_credit_usage_app_type(self._application_generate_entity.app_config.app_mode),
)
try:
+8
View File
@@ -34,6 +34,7 @@ from core.app.entities.queue_entities import (
QueueWorkflowStartedEvent,
QueueWorkflowSucceededEvent,
)
from core.credit_usage import CreditUsageAppType
from core.rag.entities import RetrievalSourceMetadata
from core.repositories.human_input_repository import HumanInputFormSubmissionRepository
from core.workflow.node_factory import (
@@ -124,6 +125,7 @@ class WorkflowBasedAppRunner:
tenant_id: str = "",
user_id: str = "",
root_node_id: str | None = None,
app_type: CreditUsageAppType | None = None,
trace_session_id: str | None = None,
) -> Graph:
"""
@@ -145,6 +147,7 @@ class WorkflowBasedAppRunner:
user_id=user_id,
user_from=user_from,
invoke_from=invoke_from,
app_type=app_type,
trace_session_id=trace_session_id,
)
graph_init_context = DifyGraphInitContext(
@@ -179,6 +182,7 @@ class WorkflowBasedAppRunner:
single_loop_run: Any | None = None,
*,
user_id: str,
app_type: CreditUsageAppType | None = None,
trace_session_id: str | None = None,
) -> tuple[Graph, VariablePool, GraphRuntimeState]:
"""
@@ -217,6 +221,7 @@ class WorkflowBasedAppRunner:
node_type_filter_key="iteration_id",
node_type_label="iteration",
user_id=user_id,
app_type=app_type,
trace_session_id=trace_session_id,
)
elif single_loop_run:
@@ -228,6 +233,7 @@ class WorkflowBasedAppRunner:
node_type_filter_key="loop_id",
node_type_label="loop",
user_id=user_id,
app_type=app_type,
trace_session_id=trace_session_id,
)
else:
@@ -247,6 +253,7 @@ class WorkflowBasedAppRunner:
node_type_label: str = "node", # 'iteration' or 'loop' for error messages
*,
user_id: str = "",
app_type: CreditUsageAppType | None = None,
trace_session_id: str | None = None,
) -> tuple[Graph, VariablePool]:
"""
@@ -313,6 +320,7 @@ class WorkflowBasedAppRunner:
user_id=user_id,
user_from=UserFrom.ACCOUNT,
invoke_from=InvokeFrom.DEBUGGER,
app_type=app_type,
trace_session_id=trace_session_id,
)
graph_init_context = DifyGraphInitContext(
@@ -6,9 +6,19 @@ from pydantic import BaseModel, ConfigDict, Field, JsonValue, ValidationInfo, fi
from constants import UUID_NIL
from core.app.app_config.entities import EasyUIBasedAppConfig, WorkflowUIBasedAppConfig
from core.credit_usage import (
CreditUsageAppType,
CreditUsageAppTypeInput,
CreditUsageCreatedBy,
CreditUsageCreatedByInput,
created_by_from_app_type,
normalize_credit_usage_app_type,
normalize_credit_usage_created_by,
)
from core.entities.provider_configuration import ProviderModelBundle
from graphon.file import File, FileUploadConfig
from graphon.model_runtime.entities.model_entities import AIModelEntity
from models.model import AppMode
if TYPE_CHECKING:
from core.ops.ops_trace_manager import TraceQueueManager
@@ -56,14 +66,58 @@ class InvokeFrom(StrEnum):
return self in (InvokeFrom.DEBUGGER, InvokeFrom.EXPLORE)
def get_credit_usage_app_type(app_mode: AppMode | str | None) -> CreditUsageAppType:
"""Return the top-level application type for an app mode."""
if app_mode is None:
return CreditUsageAppType.UNKNOWN
try:
normalized_app_mode = app_mode if isinstance(app_mode, AppMode) else AppMode.value_of(str(app_mode))
except ValueError:
return CreditUsageAppType.UNKNOWN
app_mode_mapping = {
AppMode.CHAT: CreditUsageAppType.CHATBOT,
AppMode.ADVANCED_CHAT: CreditUsageAppType.CHATFLOW,
AppMode.WORKFLOW: CreditUsageAppType.WORKFLOW,
AppMode.AGENT_CHAT: CreditUsageAppType.AGENT,
AppMode.AGENT: CreditUsageAppType.AGENT_V2,
AppMode.COMPLETION: CreditUsageAppType.COMPLETION,
AppMode.CHANNEL: CreditUsageAppType.CHANNEL,
AppMode.RAG_PIPELINE: CreditUsageAppType.RAG_PIPELINE,
}
return app_mode_mapping.get(normalized_app_mode, CreditUsageAppType.UNKNOWN)
def get_credit_usage_created_by(app_mode: AppMode | str | None) -> CreditUsageCreatedBy:
"""Return the direct app feature for an app mode."""
return created_by_from_app_type(get_credit_usage_app_type(app_mode))
class DifyRunContext(BaseModel):
tenant_id: str
app_id: str
user_id: str
user_from: UserFrom
invoke_from: InvokeFrom
app_type: CreditUsageAppType | None = None
created_by: CreditUsageCreatedBy | None = None
trace_session_id: str | None = None
@field_validator("created_by", mode="before")
@classmethod
def normalize_created_by(cls, value: object) -> CreditUsageCreatedBy | None:
if value is None:
return None
return normalize_credit_usage_created_by(value)
@field_validator("app_type", mode="before")
@classmethod
def normalize_app_type(cls, value: object) -> CreditUsageAppType | None:
if value is None:
return None
return normalize_credit_usage_app_type(value)
def build_dify_run_context(
*,
@@ -72,6 +126,8 @@ def build_dify_run_context(
user_id: str,
user_from: UserFrom,
invoke_from: InvokeFrom,
app_type: CreditUsageAppTypeInput = None,
created_by: CreditUsageCreatedByInput = None,
trace_session_id: str | None = None,
extra_context: Mapping[str, Any] | None = None,
) -> dict[str, Any]:
@@ -88,6 +144,8 @@ def build_dify_run_context(
user_id=user_id,
user_from=user_from,
invoke_from=invoke_from,
app_type=normalize_credit_usage_app_type(app_type) if app_type is not None else None,
created_by=normalize_credit_usage_created_by(created_by) if created_by is not None else None,
trace_session_id=trace_session_id,
)
return run_context
+71 -4
View File
@@ -14,6 +14,14 @@ from sqlalchemy import select
from sqlalchemy.orm import sessionmaker
from configs import dify_config
from core.credit_usage import (
CreditUsageAppType,
CreditUsageAppTypeInput,
CreditUsageCreatedBy,
CreditUsageCreatedByInput,
normalize_credit_usage_app_type,
normalize_credit_usage_created_by,
)
from core.entities.model_entities import ModelStatus
from core.entities.provider_entities import ProviderQuotaType, QuotaUnit
from core.errors.error import QuotaExceededError
@@ -25,7 +33,12 @@ 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
from services.credit_pool_service import (
CREDIT_USAGE_APP_TYPE_META_KEY,
CREDIT_USAGE_CREATED_BY_META_KEY,
CreditPoolReservation,
CreditPoolService,
)
class ModelQuotaReservationState(StrEnum):
@@ -45,6 +58,8 @@ class ModelQuotaReservation:
provider_configuration: Any
quota_unit: QuotaUnit | None = None
credit_pool_reservation: CreditPoolReservation | None = None
app_type: CreditUsageAppType | None = None
created_by: CreditUsageCreatedBy | None = None
requires_settlement: bool = False
_state: ModelQuotaReservationState = field(default=ModelQuotaReservationState.RESERVED, init=False, repr=False)
@@ -76,6 +91,10 @@ class ModelQuotaReservation:
provider=self.provider,
provider_configuration=self.provider_configuration,
used_quota=used_quota,
model_type=self.model_type,
model=self.model,
app_type=self.app_type,
created_by=self.created_by,
)
self._state = ModelQuotaReservationState.COMMITTED
@@ -121,15 +140,21 @@ def reserve_model_quota_for_model(
model_type: ModelType,
model: str,
request_id: str | None = None,
app_type: CreditUsageAppTypeInput = None,
created_by: CreditUsageCreatedByInput = None,
) -> ModelQuotaReservation:
"""Reserve system-hosted model quota before invoking the provider."""
provider_configuration = _get_provider_configuration(tenant_id=tenant_id, provider=provider)
effective_app_type = normalize_credit_usage_app_type(app_type)
effective_created_by = normalize_credit_usage_created_by(created_by)
reservation = ModelQuotaReservation(
tenant_id=tenant_id,
provider=provider,
model_type=model_type,
model=model,
provider_configuration=provider_configuration,
app_type=effective_app_type,
created_by=effective_created_by,
)
if provider_configuration.using_provider_type != ProviderType.SYSTEM:
return reservation
@@ -166,6 +191,8 @@ def reserve_model_quota_for_model(
"model_type": model_type.value,
"model": model,
}
reservation_meta[CREDIT_USAGE_CREATED_BY_META_KEY] = effective_created_by
reservation_meta[CREDIT_USAGE_APP_TYPE_META_KEY] = effective_app_type
reservation.credit_pool_reservation = CreditPoolService.reserve_credits(
tenant_id=tenant_id,
credits_required=amount,
@@ -183,7 +210,13 @@ def reserve_model_quota_for_model(
def reserve_llm_quota_for_model(
*, tenant_id: str, provider: str, model: str, request_id: str | None = None
*,
tenant_id: str,
provider: str,
model: str,
request_id: str | None = None,
app_type: CreditUsageAppTypeInput = None,
created_by: CreditUsageCreatedByInput = None,
) -> ModelQuotaReservation:
"""Reserve system-hosted LLM quota before invoking the provider."""
return reserve_model_quota_for_model(
@@ -192,6 +225,8 @@ def reserve_llm_quota_for_model(
model_type=ModelType.LLM,
model=model,
request_id=request_id,
app_type=app_type,
created_by=created_by,
)
@@ -286,13 +321,31 @@ def _deduct_free_model_quota(
raise QuotaExceededError(f"Model provider {provider} quota exceeded.")
def _deduct_used_model_quota(*, tenant_id: str, provider: str, provider_configuration, used_quota: int | None) -> None:
def _deduct_used_model_quota(
*,
tenant_id: str,
provider: str,
provider_configuration,
used_quota: int | None,
model_type: ModelType | None = None,
model: str | None = None,
app_type: CreditUsageAppTypeInput = None,
created_by: CreditUsageCreatedByInput = None,
) -> None:
"""Apply a resolved model quota charge against the current provider quota bucket."""
if provider_configuration.using_provider_type != ProviderType.SYSTEM:
return
system_configuration = provider_configuration.system_configuration
if used_quota is not None and system_configuration.current_quota_type is not None:
metadata: dict[str, object] = {"provider": provider}
if model is not None:
metadata["model"] = model
if model_type is not None:
metadata["model_type"] = model_type.value
metadata[CREDIT_USAGE_APP_TYPE_META_KEY] = normalize_credit_usage_app_type(app_type)
metadata[CREDIT_USAGE_CREATED_BY_META_KEY] = normalize_credit_usage_created_by(created_by)
match system_configuration.current_quota_type:
case ProviderQuotaType.TRIAL:
from services.credit_pool_service import CreditPoolService
@@ -300,6 +353,7 @@ def _deduct_used_model_quota(*, tenant_id: str, provider: str, provider_configur
CreditPoolService.deduct_credits_capped(
tenant_id=tenant_id,
credits_required=used_quota,
metadata=metadata,
session=db.session(),
)
case ProviderQuotaType.PAID:
@@ -309,6 +363,7 @@ def _deduct_used_model_quota(*, tenant_id: str, provider: str, provider_configur
tenant_id=tenant_id,
credits_required=used_quota,
pool_type="paid",
metadata=metadata,
session=db.session(),
)
case ProviderQuotaType.FREE:
@@ -322,7 +377,15 @@ def _deduct_used_model_quota(*, tenant_id: str, provider: str, provider_configur
return
def deduct_llm_quota_for_model(*, tenant_id: str, provider: str, model: str, usage: LLMUsage) -> None:
def deduct_llm_quota_for_model(
*,
tenant_id: str,
provider: str,
model: str,
usage: LLMUsage,
app_type: CreditUsageAppTypeInput = None,
created_by: CreditUsageCreatedByInput = None,
) -> None:
"""Deduct tenant-bound quota for the resolved LLM model identity."""
provider_configuration = _get_provider_configuration(tenant_id=tenant_id, provider=provider)
used_quota = _resolve_llm_used_quota(
@@ -335,6 +398,10 @@ def deduct_llm_quota_for_model(*, tenant_id: str, provider: str, model: str, usa
provider=provider,
provider_configuration=provider_configuration,
used_quota=used_quota,
model_type=ModelType.LLM,
model=model,
app_type=app_type,
created_by=created_by,
)
@@ -13,6 +13,7 @@ from core.app.entities.app_invoke_entities import (
AgentChatAppGenerateEntity,
ChatAppGenerateEntity,
CompletionAppGenerateEntity,
get_credit_usage_app_type,
)
from core.app.entities.queue_entities import (
QueueAgentMessageEvent,
@@ -231,7 +232,10 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline[EasyUIAppGenerat
and text_to_speech_dict.get("enabled")
):
publisher = AppGeneratorTTSPublisher(
tenant_id, text_to_speech_dict.get("voice", ""), text_to_speech_dict.get("language", None)
tenant_id,
text_to_speech_dict.get("voice", ""),
text_to_speech_dict.get("language", None),
get_credit_usage_app_type(self._app_config.app_mode),
)
try:
for response in self._process_stream_response(publisher=publisher, trace_manager=trace_manager):
@@ -17,6 +17,7 @@ from core.base.tts.audio_mime import (
get_model_audio_mime_type,
inspect_audio_stream,
)
from core.credit_usage import CreditUsageAppType, CreditUsageCreatedBy
from core.model_manager import ModelManager
from graphon.model_runtime.entities.message_entities import TextPromptMessageContent
from graphon.model_runtime.entities.model_entities import ModelType
@@ -32,14 +33,25 @@ class AudioTrunk:
class AppGeneratorTTSPublisher:
def __init__(self, tenant_id: str, voice: str, language: str | None = None):
def __init__(
self,
tenant_id: str,
voice: str,
language: str | None = None,
app_type: CreditUsageAppType = CreditUsageAppType.UNKNOWN,
created_by: CreditUsageCreatedBy = CreditUsageCreatedBy.AUDIO,
):
self.logger = logging.getLogger(__name__)
self.tenant_id = tenant_id
self.msg_text = ""
self._audio_queue: queue.Queue[AudioTrunk] = queue.Queue()
self._msg_queue: queue.Queue[WorkflowQueueMessage | MessageQueueMessage | None] = queue.Queue()
self.match = re.compile(r"[。.!?]")
self.model_manager = ModelManager.for_tenant(tenant_id=self.tenant_id, user_id="responding_tts")
self.model_manager = ModelManager.for_tenant(
tenant_id=self.tenant_id,
user_id="responding_tts",
request_metadata={"app_type": app_type, "created_by": created_by},
)
self.model_instance = self.model_manager.get_default_model_instance(
tenant_id=self.tenant_id, model_type=ModelType.TTS
)
+74
View File
@@ -0,0 +1,74 @@
from enum import StrEnum
class CreditUsageCreatedBy(StrEnum):
"""Business feature that created a credit usage event."""
APP = "app"
AGENT_NODE = "agent_node"
BUILD_DRAFT = "build_draft"
CONVERSATION_NAME = "conversation_name"
SUGGESTED_QUESTIONS = "suggested_questions"
WORKFLOW_GENERATION = "workflow_generation"
WORKFLOW_INSTRUCTION_SUGGESTIONS = "workflow_instruction_suggestions"
RULE_CONFIG = "rule_config"
CODE_GENERATION = "code_generation"
QA_DOCUMENT = "qa_document"
STRUCTURED_OUTPUT = "structured_output"
INSTRUCTION_MODIFICATION = "instruction_modification"
KNOWLEDGE_RETRIEVAL = "knowledge_retrieval"
KNOWLEDGE_INDEXING = "knowledge_indexing"
TOOL = "tool"
AUDIO = "audio"
MODERATION = "moderation"
PLUGIN_API = "plugin_api"
UNKNOWN = "unknown"
type CreditUsageCreatedByInput = CreditUsageCreatedBy | str | None
class CreditUsageAppType(StrEnum):
"""Top-level application type associated with a credit usage event."""
CHATBOT = "chatbot"
CHATFLOW = "chatflow"
WORKFLOW = "workflow"
AGENT = "agent"
AGENT_V2 = "agent_v2"
COMPLETION = "completion"
CHANNEL = "channel"
RAG_PIPELINE = "rag_pipeline"
UNKNOWN = "unknown"
type CreditUsageAppTypeInput = CreditUsageAppType | str | None
def normalize_credit_usage_created_by(value: object) -> CreditUsageCreatedBy:
if isinstance(value, CreditUsageCreatedBy):
return value
if isinstance(value, str):
try:
return CreditUsageCreatedBy(value)
except ValueError:
pass
return CreditUsageCreatedBy.UNKNOWN
def normalize_credit_usage_app_type(value: object) -> CreditUsageAppType:
if isinstance(value, CreditUsageAppType):
return value
if isinstance(value, str):
try:
return CreditUsageAppType(value)
except ValueError:
pass
return CreditUsageAppType.UNKNOWN
def created_by_from_app_type(app_type: CreditUsageAppTypeInput) -> CreditUsageCreatedBy:
normalized_app_type = normalize_credit_usage_app_type(app_type)
if normalized_app_type is CreditUsageAppType.UNKNOWN:
return CreditUsageCreatedBy.UNKNOWN
return CreditUsageCreatedBy.APP
+9 -2
View File
@@ -14,9 +14,11 @@ from sqlalchemy.orm import Session
from sqlalchemy.orm.exc import ObjectDeletedError
from configs import dify_config
from core.credit_usage import CreditUsageCreatedBy
from core.db.session_factory import session_factory
from core.entities.knowledge_entities import IndexingEstimate, PreviewDetail, QAPreviewDetail
from core.errors.error import ProviderTokenNotInitError
from core.model_context import with_credit_usage_created_by
from core.model_manager import ModelInstance, ModelManager
from core.rag.cleaner.clean_processor import CleanProcessor
from core.rag.datasource.keyword.keyword_factory import Keyword
@@ -37,6 +39,7 @@ from core.tools.utils.web_reader_tool import get_image_upload_file_ids
from enums import DeploymentEdition
from extensions.ext_redis import redis_client
from extensions.ext_storage import storage
from extensions.otel import propagate_context
from graphon.model_runtime.entities.model_entities import ModelType
from libs import helper
from libs.datetime_utils import naive_utc_now
@@ -74,6 +77,7 @@ class IndexingRunner:
document.stopped_at = naive_utc_now()
session.flush()
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_INDEXING)
def run(self, dataset_documents: list[DatasetDocument], session: Session):
"""Run indexing with commits before slow transforms and parallel index workers.
@@ -160,6 +164,7 @@ class IndexingRunner:
except Exception as e:
self._handle_indexing_error(document_id, e, session)
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_INDEXING)
def run_in_splitting_status(self, dataset_document: DatasetDocument, session: Session):
"""Run the indexing process when the index_status is splitting."""
document_id = dataset_document.id
@@ -243,6 +248,7 @@ class IndexingRunner:
except Exception as e:
self._handle_indexing_error(document_id, e, session)
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_INDEXING)
def run_in_indexing_status(self, dataset_document: DatasetDocument, session: Session):
"""Run the indexing process when the index_status is indexing."""
document_id = dataset_document.id
@@ -315,6 +321,7 @@ class IndexingRunner:
except Exception as e:
self._handle_indexing_error(document_id, e, session)
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_INDEXING)
def indexing_estimate(
self,
tenant_id: str,
@@ -652,7 +659,7 @@ class IndexingRunner:
):
# create keyword index
create_keyword_thread = threading.Thread(
target=self._process_keyword_index,
target=propagate_context(self._process_keyword_index),
args=(current_app._get_current_object(), dataset.id, dataset_document.id, documents), # type: ignore
)
create_keyword_thread.start()
@@ -675,7 +682,7 @@ class IndexingRunner:
continue
futures.append(
executor.submit(
self._process_chunk,
propagate_context(self._process_chunk),
current_app._get_current_object(), # type: ignore
dataset_document.doc_form,
chunk_documents,
+11
View File
@@ -9,6 +9,7 @@ from sqlalchemy import select
from sqlalchemy.orm import Session
from core.app.app_config.entities import ModelConfig
from core.app.entities.app_invoke_entities import CreditUsageCreatedBy
from core.llm_generator.entities import RuleCodeGeneratePayload, RuleGeneratePayload, RuleStructuredOutputPayload
from core.llm_generator.output_parser.rule_config_generator import RuleConfigGeneratorOutputParser
from core.llm_generator.output_parser.suggested_questions_after_answer import SuggestedQuestionsAfterAnswerOutputParser
@@ -22,6 +23,7 @@ from core.llm_generator.prompts import (
SYSTEM_STRUCTURED_OUTPUT_GENERATE,
WORKFLOW_RULE_CONFIG_PROMPT_GENERATE_TEMPLATE,
)
from core.model_context import with_credit_usage_created_by
from core.model_manager import ModelManager
from core.ops.entities.trace_entity import TraceTaskName
from core.ops.ops_trace_manager import TraceQueueManager, TraceTask
@@ -187,6 +189,7 @@ class LLMGenerator:
logger.debug("Failed to emit prompt_generation telemetry", exc_info=True)
@classmethod
@with_credit_usage_created_by(CreditUsageCreatedBy.CONVERSATION_NAME)
def generate_conversation_name(
cls,
tenant_id: str,
@@ -255,6 +258,7 @@ class LLMGenerator:
return name
@classmethod
@with_credit_usage_created_by(CreditUsageCreatedBy.SUGGESTED_QUESTIONS)
def generate_suggested_questions_after_answer(
cls,
tenant_id: str,
@@ -342,6 +346,7 @@ class LLMGenerator:
return questions
@classmethod
@with_credit_usage_created_by(CreditUsageCreatedBy.WORKFLOW_INSTRUCTION_SUGGESTIONS)
def generate_workflow_instruction_suggestions(
cls,
tenant_id: str,
@@ -457,6 +462,7 @@ class LLMGenerator:
return "\n\n".join(sections) + "\n\n"
@classmethod
@with_credit_usage_created_by(CreditUsageCreatedBy.RULE_CONFIG)
def generate_rule_config(cls, tenant_id: str, args: RuleGeneratePayload, *, app_id: str | None = None):
output_parser = RuleConfigGeneratorOutputParser()
@@ -624,6 +630,7 @@ class LLMGenerator:
return rule_config
@classmethod
@with_credit_usage_created_by(CreditUsageCreatedBy.CODE_GENERATION)
def generate_code(
cls,
tenant_id: str,
@@ -692,6 +699,7 @@ class LLMGenerator:
return {"code": generated_code, "language": args.code_language, "error": ""}
@classmethod
@with_credit_usage_created_by(CreditUsageCreatedBy.QA_DOCUMENT)
def generate_qa_document(cls, tenant_id: str, query, document_language: str):
prompt = GENERATOR_QA_PROMPT.format(language=document_language)
@@ -719,6 +727,7 @@ class LLMGenerator:
return answer.strip()
@classmethod
@with_credit_usage_created_by(CreditUsageCreatedBy.STRUCTURED_OUTPUT)
def generate_structured_output(
cls, tenant_id: str, args: RuleStructuredOutputPayload, *, app_id: str | None = None
) -> StructuredOutputResultDict:
@@ -776,6 +785,7 @@ class LLMGenerator:
return {"output": generated_output, "error": ""}
@staticmethod
@with_credit_usage_created_by(CreditUsageCreatedBy.INSTRUCTION_MODIFICATION)
def instruction_modify_legacy(
tenant_id: str,
flow_id: str,
@@ -821,6 +831,7 @@ class LLMGenerator:
)
@staticmethod
@with_credit_usage_created_by(CreditUsageCreatedBy.INSTRUCTION_MODIFICATION)
def instruction_modify_workflow(
tenant_id: str,
flow_id: str,
@@ -56,6 +56,7 @@ def invoke_llm_with_structured_output(
stop: list[str] | None = None,
stream: Literal[True],
callbacks: list[Callback] | None = None,
request_metadata: Mapping[str, object] | None = None,
) -> Generator[LLMResultChunkWithStructuredOutput, None, None]: ...
@overload
def invoke_llm_with_structured_output(
@@ -70,6 +71,7 @@ def invoke_llm_with_structured_output(
stop: list[str] | None = None,
stream: Literal[False],
callbacks: list[Callback] | None = None,
request_metadata: Mapping[str, object] | None = None,
) -> LLMResultWithStructuredOutput: ...
@overload
def invoke_llm_with_structured_output(
@@ -84,6 +86,7 @@ def invoke_llm_with_structured_output(
stop: list[str] | None = None,
stream: bool = True,
callbacks: list[Callback] | None = None,
request_metadata: Mapping[str, object] | None = None,
) -> LLMResultWithStructuredOutput | Generator[LLMResultChunkWithStructuredOutput, None, None]: ...
def invoke_llm_with_structured_output(
*,
@@ -97,6 +100,7 @@ def invoke_llm_with_structured_output(
stop: list[str] | None = None,
stream: bool = True,
callbacks: list[Callback] | None = None,
request_metadata: Mapping[str, object] | None = None,
) -> LLMResultWithStructuredOutput | Generator[LLMResultChunkWithStructuredOutput, None, None]:
"""
Invoke large language model with structured output
@@ -139,6 +143,7 @@ def invoke_llm_with_structured_output(
stop=stop,
stream=stream,
callbacks=callbacks,
request_metadata=request_metadata,
)
if isinstance(llm_result, LLMResult):
+61
View File
@@ -0,0 +1,61 @@
from collections.abc import Callable, Generator, Mapping
from contextlib import contextmanager
from contextvars import ContextVar
from functools import wraps
from typing import Any, Protocol, cast
from core.credit_usage import CreditUsageCreatedByInput, normalize_credit_usage_created_by
_credit_usage_metadata: ContextVar[dict[str, object] | None] = ContextVar(
"credit_usage_metadata",
default=None,
)
class _CreditUsageMetadataCarrier(Protocol):
_request_metadata: Mapping[str, object] | None
def get_credit_usage_metadata() -> Mapping[str, object] | None:
return _credit_usage_metadata.get()
@contextmanager
def use_credit_usage_metadata(metadata: Mapping[str, object] | None) -> Generator[None, None, None]:
if metadata is None:
yield
return
current_metadata = _credit_usage_metadata.get()
effective_metadata = {**metadata, **(current_metadata or {})}
token = _credit_usage_metadata.set(effective_metadata)
try:
yield
finally:
_credit_usage_metadata.reset(token)
def with_credit_usage_metadata(method: Callable[..., Any]) -> Callable[..., Any]:
@wraps(method)
def wrapper(self, *args: object, **kwargs: object):
carrier = cast(_CreditUsageMetadataCarrier, self)
with use_credit_usage_metadata(carrier._request_metadata):
return method(self, *args, **kwargs)
return wrapper
def with_credit_usage_created_by(
created_by: CreditUsageCreatedByInput,
) -> Callable[[Callable[..., Any]], Callable[..., Any]]:
normalized_created_by = normalize_credit_usage_created_by(created_by)
def decorator(method: Callable[..., Any]) -> Callable[..., Any]:
@wraps(method)
def wrapper(*args: object, **kwargs: object):
with use_credit_usage_metadata({"created_by": normalized_created_by}):
return method(*args, **kwargs)
return wrapper
return decorator
+185 -27
View File
@@ -5,11 +5,20 @@ from typing import IO, Any, Literal, Optional, ParamSpec, TypeVar, Union, cast,
from uuid import UUID
from configs import dify_config
from core.credit_usage import (
CreditUsageAppType,
CreditUsageAppTypeInput,
CreditUsageCreatedBy,
CreditUsageCreatedByInput,
normalize_credit_usage_app_type,
normalize_credit_usage_created_by,
)
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 ModelCurrentlyNotSupportError, ProviderTokenNotInitError
from core.model_context import get_credit_usage_metadata
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
@@ -38,13 +47,20 @@ class ModelInstance:
Model instance class.
"""
def __init__(self, provider_model_bundle: ProviderModelBundle, model: str, credentials: dict | None = None) -> None:
def __init__(
self,
provider_model_bundle: ProviderModelBundle,
model: str,
credentials: dict | None = None,
request_metadata: Mapping[str, object] | None = None,
) -> None:
self.provider_model_bundle = provider_model_bundle
self.model_name = model
self.provider = provider_model_bundle.configuration.provider.provider
if credentials is None:
credentials = self._fetch_credentials_from_bundle(provider_model_bundle, model)
self.credentials = credentials
self._request_metadata = dict(request_metadata) if request_metadata else None
# Runtime LLM invocation fields.
self.parameters: Mapping[str, Any] = {}
self.stop: Sequence[str] = ()
@@ -56,6 +72,16 @@ class ModelInstance:
credentials=self.credentials,
)
def _resolve_request_metadata(self, request_metadata: Mapping[str, object] | None) -> Mapping[str, object] | None:
bound_request_metadata = self._request_metadata
if bound_request_metadata is None:
if request_metadata is None:
return get_credit_usage_metadata()
return request_metadata
if request_metadata is None:
return bound_request_metadata
return {**bound_request_metadata, **request_metadata}
def get_model_schema(self) -> AIModelEntity:
"""Return the resolved schema for the current model instance."""
model_schema = self.model_type_instance.get_model_schema(self.model_name, self.credentials)
@@ -176,6 +202,7 @@ class ModelInstance:
"""
if not isinstance(self.model_type_instance, LargeLanguageModel):
raise Exception("Model type instance is not LargeLanguageModel")
request_metadata = self._resolve_request_metadata(request_metadata)
return cast(
Union[LLMResult, Generator],
self._round_robin_invoke(
@@ -213,7 +240,11 @@ class ModelInstance:
)
def invoke_text_embedding(
self, texts: list[str], input_type: EmbeddingInputType = EmbeddingInputType.DOCUMENT
self,
texts: list[str],
input_type: EmbeddingInputType = EmbeddingInputType.DOCUMENT,
*,
request_metadata: Mapping[str, object] | None = None,
) -> EmbeddingResult:
"""
Invoke large language model
@@ -230,12 +261,15 @@ class ModelInstance:
credentials=self.credentials,
texts=texts,
input_type=input_type,
request_metadata=self._resolve_request_metadata(request_metadata),
)
def invoke_multimodal_embedding(
self,
multimodel_documents: list[dict],
input_type: EmbeddingInputType = EmbeddingInputType.DOCUMENT,
*,
request_metadata: Mapping[str, object] | None = None,
) -> EmbeddingResult:
"""
Invoke large language model
@@ -252,6 +286,7 @@ class ModelInstance:
credentials=self.credentials,
multimodel_documents=multimodel_documents,
input_type=input_type,
request_metadata=self._resolve_request_metadata(request_metadata),
)
def get_text_embedding_num_tokens(self, texts: list[str]) -> list[int]:
@@ -276,6 +311,8 @@ class ModelInstance:
docs: list[str],
score_threshold: float | None = None,
top_n: int | None = None,
*,
request_metadata: Mapping[str, object] | None = None,
) -> RerankResult:
"""
Invoke rerank model
@@ -296,6 +333,7 @@ class ModelInstance:
docs=docs,
score_threshold=score_threshold,
top_n=top_n,
request_metadata=self._resolve_request_metadata(request_metadata),
)
def invoke_multimodal_rerank(
@@ -304,6 +342,8 @@ class ModelInstance:
docs: list[MultimodalRerankInput],
score_threshold: float | None = None,
top_n: int | None = None,
*,
request_metadata: Mapping[str, object] | None = None,
) -> RerankResult:
"""
Invoke rerank model
@@ -324,9 +364,10 @@ class ModelInstance:
docs=docs,
score_threshold=score_threshold,
top_n=top_n,
request_metadata=self._resolve_request_metadata(request_metadata),
)
def invoke_moderation(self, text: str) -> bool:
def invoke_moderation(self, text: str, *, request_metadata: Mapping[str, object] | None = None) -> bool:
"""
Invoke moderation model
@@ -340,9 +381,10 @@ class ModelInstance:
model=self.model_name,
credentials=self.credentials,
text=text,
request_metadata=self._resolve_request_metadata(request_metadata),
)
def invoke_speech2text(self, file: IO[bytes]) -> str:
def invoke_speech2text(self, file: IO[bytes], *, request_metadata: Mapping[str, object] | None = None) -> str:
"""
Invoke large language model
@@ -356,9 +398,16 @@ class ModelInstance:
model=self.model_name,
credentials=self.credentials,
file=file,
request_metadata=self._resolve_request_metadata(request_metadata),
)
def invoke_tts(self, content_text: str, voice: str = "") -> Iterable[bytes]:
def invoke_tts(
self,
content_text: str,
voice: str = "",
*,
request_metadata: Mapping[str, object] | None = None,
) -> Iterable[bytes]:
"""
Invoke large language tts model
@@ -374,6 +423,7 @@ class ModelInstance:
credentials=self.credentials,
content_text=content_text,
voice=voice,
request_metadata=self._resolve_request_metadata(request_metadata),
)
def _round_robin_invoke(self, function: Callable[P, R], *args: P.args, **kwargs: P.kwargs) -> R:
@@ -446,15 +496,28 @@ class ModelInstance:
class QuotaManagedModelInstance(ModelInstance):
"""A system-hosted model instance that owns quota settlement per invocation."""
def reserve_quota(self, *, request_id: str | None = None):
def reserve_quota(
self,
*,
request_id: str | None = None,
app_type: CreditUsageAppTypeInput = None,
created_by: CreditUsageCreatedByInput = None,
):
from core.app.llm.quota import reserve_model_quota_for_model
if app_type is None:
app_type = self._get_reservation_app_type(self._request_metadata)
if created_by is None:
created_by = self._get_reservation_created_by(self._request_metadata)
return reserve_model_quota_for_model(
tenant_id=self.provider_model_bundle.configuration.tenant_id,
provider=self.provider,
model_type=self.model_type_instance.model_type,
model=self.model_name,
request_id=request_id,
app_type=app_type,
created_by=created_by,
)
@staticmethod
@@ -467,11 +530,39 @@ class QuotaManagedModelInstance(ModelInstance):
except ValueError:
return None
@staticmethod
def _get_reservation_created_by(
request_metadata: Mapping[str, object] | None,
) -> CreditUsageCreatedBy | None:
created_by = request_metadata.get("created_by") if request_metadata else None
if created_by is None:
return None
return normalize_credit_usage_created_by(created_by)
@staticmethod
def _get_reservation_app_type(
request_metadata: Mapping[str, object] | None,
) -> CreditUsageAppType | None:
app_type = request_metadata.get("app_type") if request_metadata else None
if app_type is None:
return None
return normalize_credit_usage_app_type(app_type)
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:
app_type = self._get_reservation_app_type(request_metadata)
created_by = self._get_reservation_created_by(request_metadata)
if request_id is None and app_type is None and created_by is None:
return self.reserve_quota()
return self.reserve_quota(request_id=request_id)
reservation_kwargs: dict[str, str | CreditUsageAppType | CreditUsageCreatedBy] = {}
if request_id is not None:
reservation_kwargs["request_id"] = request_id
if app_type is not None:
reservation_kwargs["app_type"] = app_type
if created_by is not None:
reservation_kwargs["created_by"] = created_by
return self.reserve_quota(**reservation_kwargs)
@staticmethod
def release_quota_safely(reservation) -> None:
@@ -481,7 +572,13 @@ class QuotaManagedModelInstance(ModelInstance):
logger.exception("Failed to release model quota reservation")
def _invoke_with_quota(self, function: Callable[P, R], *args: P.args, **kwargs: P.kwargs) -> R:
reservation = self.reserve_quota()
request_metadata = kwargs.get("request_metadata")
effective_request_metadata = self._resolve_request_metadata(
request_metadata if isinstance(request_metadata, Mapping) else None
)
if effective_request_metadata is not None:
kwargs["request_metadata"] = effective_request_metadata
reservation = self._reserve_quota_for_request(effective_request_metadata)
try:
response = function(*args, **kwargs)
reservation.commit()
@@ -548,7 +645,8 @@ class QuotaManagedModelInstance(ModelInstance):
request_metadata=request_metadata,
)
reservation = self._reserve_quota_for_request(request_metadata)
effective_request_metadata = self._resolve_request_metadata(request_metadata)
reservation = self._reserve_quota_for_request(effective_request_metadata)
try:
response = super().invoke_llm(
prompt_messages=normalized_prompt_messages,
@@ -557,7 +655,7 @@ class QuotaManagedModelInstance(ModelInstance):
stop=normalized_stop,
stream=False,
callbacks=callbacks,
request_metadata=request_metadata,
request_metadata=effective_request_metadata,
)
if isinstance(response, Generator):
raise TypeError("Non-streaming LLM invocation returned a generator.")
@@ -576,7 +674,8 @@ class QuotaManagedModelInstance(ModelInstance):
callbacks: list[Callback] | None,
request_metadata: Mapping[str, object] | None,
) -> Generator:
reservation = self._reserve_quota_for_request(request_metadata)
effective_request_metadata = self._resolve_request_metadata(request_metadata)
reservation = self._reserve_quota_for_request(effective_request_metadata)
usage: LLMUsage | None = None
try:
response = super().invoke_llm(
@@ -586,7 +685,7 @@ class QuotaManagedModelInstance(ModelInstance):
stop=stop,
stream=True,
callbacks=callbacks,
request_metadata=request_metadata,
request_metadata=effective_request_metadata,
)
if not isinstance(response, Generator):
raise TypeError("Streaming LLM invocation did not return a generator.")
@@ -614,20 +713,32 @@ class QuotaManagedModelInstance(ModelInstance):
@override
def invoke_text_embedding(
self, texts: list[str], input_type: EmbeddingInputType = EmbeddingInputType.DOCUMENT
self,
texts: list[str],
input_type: EmbeddingInputType = EmbeddingInputType.DOCUMENT,
*,
request_metadata: Mapping[str, object] | None = None,
) -> EmbeddingResult:
return self._invoke_with_quota(super().invoke_text_embedding, texts=texts, input_type=input_type)
return self._invoke_with_quota(
super().invoke_text_embedding,
texts=texts,
input_type=input_type,
request_metadata=request_metadata,
)
@override
def invoke_multimodal_embedding(
self,
multimodel_documents: list[dict],
input_type: EmbeddingInputType = EmbeddingInputType.DOCUMENT,
*,
request_metadata: Mapping[str, object] | None = None,
) -> EmbeddingResult:
return self._invoke_with_quota(
super().invoke_multimodal_embedding,
multimodel_documents=multimodel_documents,
input_type=input_type,
request_metadata=request_metadata,
)
@override
@@ -637,6 +748,8 @@ class QuotaManagedModelInstance(ModelInstance):
docs: list[str],
score_threshold: float | None = None,
top_n: int | None = None,
*,
request_metadata: Mapping[str, object] | None = None,
) -> RerankResult:
return self._invoke_with_quota(
super().invoke_rerank,
@@ -644,6 +757,7 @@ class QuotaManagedModelInstance(ModelInstance):
docs=docs,
score_threshold=score_threshold,
top_n=top_n,
request_metadata=request_metadata,
)
@override
@@ -653,6 +767,8 @@ class QuotaManagedModelInstance(ModelInstance):
docs: list[MultimodalRerankInput],
score_threshold: float | None = None,
top_n: int | None = None,
*,
request_metadata: Mapping[str, object] | None = None,
) -> RerankResult:
return self._invoke_with_quota(
super().invoke_multimodal_rerank,
@@ -660,24 +776,42 @@ class QuotaManagedModelInstance(ModelInstance):
docs=docs,
score_threshold=score_threshold,
top_n=top_n,
request_metadata=request_metadata,
)
@override
def invoke_moderation(self, text: str) -> bool:
return self._invoke_with_quota(super().invoke_moderation, text=text)
def invoke_moderation(self, text: str, *, request_metadata: Mapping[str, object] | None = None) -> bool:
return self._invoke_with_quota(super().invoke_moderation, text=text, request_metadata=request_metadata)
@override
def invoke_speech2text(self, file: IO[bytes]) -> str:
return self._invoke_with_quota(super().invoke_speech2text, file=file)
def invoke_speech2text(self, file: IO[bytes], *, request_metadata: Mapping[str, object] | None = None) -> str:
return self._invoke_with_quota(super().invoke_speech2text, file=file, request_metadata=request_metadata)
@override
def invoke_tts(self, content_text: str, voice: str = "") -> Iterable[bytes]:
return self._invoke_tts_stream(content_text=content_text, voice=voice)
def invoke_tts(
self,
content_text: str,
voice: str = "",
*,
request_metadata: Mapping[str, object] | None = None,
) -> Iterable[bytes]:
return self._invoke_tts_stream(content_text=content_text, voice=voice, request_metadata=request_metadata)
def _invoke_tts_stream(self, *, content_text: str, voice: str) -> Generator[bytes, None, None]:
reservation = self.reserve_quota()
def _invoke_tts_stream(
self,
*,
content_text: str,
voice: str,
request_metadata: Mapping[str, object] | None = None,
) -> Generator[bytes, None, None]:
effective_request_metadata = self._resolve_request_metadata(request_metadata)
reservation = self._reserve_quota_for_request(effective_request_metadata)
try:
response = super().invoke_tts(content_text=content_text, voice=voice)
response = super().invoke_tts(
content_text=content_text,
voice=voice,
request_metadata=effective_request_metadata,
)
for chunk in response:
reservation.commit()
yield chunk
@@ -706,14 +840,34 @@ class ModelManager:
provider_manager: ProviderManager,
*,
enable_credentials_cache: bool = False,
request_metadata: Mapping[str, object] | None = None,
) -> None:
self._provider_manager = provider_manager
self._credentials_cache: dict[tuple[str, str, str, str], Any] = {}
self._enable_credentials_cache = enable_credentials_cache
self._request_metadata = dict(request_metadata) if request_metadata else None
@classmethod
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))
def for_tenant(
cls,
tenant_id: str,
user_id: str | None = None,
*,
request_metadata: Mapping[str, object] | None = None,
) -> "ModelManager":
if request_metadata is None:
request_metadata = get_credit_usage_metadata()
return cls(
provider_manager=create_plugin_provider_manager(tenant_id=tenant_id, user_id=user_id),
request_metadata=request_metadata,
)
def _resolve_request_metadata(self, request_metadata: Mapping[str, object] | None) -> Mapping[str, object] | None:
if self._request_metadata is None:
return request_metadata
if request_metadata is None:
return self._request_metadata
return {**self._request_metadata, **request_metadata}
@staticmethod
def _validate_system_model_access(
@@ -755,6 +909,8 @@ class ModelManager:
provider: str,
model_type: ModelType,
model: str,
*,
request_metadata: Mapping[str, object] | None = None,
) -> ModelInstance:
"""
Get model instance
@@ -772,6 +928,7 @@ class ModelManager:
)
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)
effective_request_metadata = self._resolve_request_metadata(request_metadata)
cred_cache_key = (tenant_id, provider, model_type.value, model)
@@ -780,9 +937,10 @@ class ModelManager:
provider_model_bundle,
model,
deepcopy(self._credentials_cache[cred_cache_key]),
effective_request_metadata,
)
ret = model_instance_class(provider_model_bundle, model)
ret = model_instance_class(provider_model_bundle, model, request_metadata=effective_request_metadata)
if self._enable_credentials_cache:
self._credentials_cache[cred_cache_key] = deepcopy(ret.credentials)
return ret
+8 -5
View File
@@ -3,6 +3,8 @@ from collections.abc import Mapping
from typing import Any
from core.app.app_config.entities import AppConfig
from core.app.entities.app_invoke_entities import get_credit_usage_app_type
from core.model_context import use_credit_usage_metadata
from core.moderation.base import ModerationAction, ModerationError
from core.moderation.factory import ModerationFactory
from core.ops.entities.trace_entity import TraceTaskName
@@ -41,12 +43,13 @@ class InputModeration:
sensitive_word_avoidance_config = app_config.sensitive_word_avoidance
moderation_type = sensitive_word_avoidance_config.type
moderation_factory = ModerationFactory(
name=moderation_type, app_id=app_id, tenant_id=tenant_id, config=sensitive_word_avoidance_config.config
)
with use_credit_usage_metadata({"app_type": get_credit_usage_app_type(app_config.app_mode)}):
moderation_factory = ModerationFactory(
name=moderation_type, app_id=app_id, tenant_id=tenant_id, config=sensitive_word_avoidance_config.config
)
with measure_time() as timer:
moderation_result = moderation_factory.moderation_for_inputs(inputs, query)
with measure_time() as timer:
moderation_result = moderation_factory.moderation_for_inputs(inputs, query)
if trace_manager:
trace_manager.add_trace_task(
@@ -1,5 +1,7 @@
from typing import Any, override
from core.credit_usage import CreditUsageCreatedBy
from core.model_context import with_credit_usage_created_by
from core.model_manager import ModelManager
from core.moderation.base import Moderation, ModerationAction, ModerationInputsResult, ModerationOutputsResult
from graphon.model_runtime.entities.model_entities import ModelType
@@ -53,6 +55,7 @@ class OpenAIModeration(Moderation):
flagged=flagged, action=ModerationAction.DIRECT_OUTPUT, preset_response=preset_response
)
@with_credit_usage_created_by(CreditUsageCreatedBy.MODERATION)
def _is_violated(self, inputs: dict[str, Any]):
text = "\n".join(str(inputs.values()))
model_manager = ModelManager.for_tenant(tenant_id=self.tenant_id)
@@ -4,7 +4,9 @@ from collections.abc import Generator
from typing import Any
from core.base.tts.audio_mime import get_model_audio_mime_type, inspect_audio_stream
from core.credit_usage import CreditUsageCreatedBy
from core.llm_generator.output_parser.structured_output import invoke_llm_with_structured_output
from core.model_context import with_credit_usage_created_by
from core.model_manager import ModelManager
from core.plugin.backwards_invocation.base import BaseBackwardsInvocation
from core.plugin.entities.request import (
@@ -37,6 +39,7 @@ from models.account import Tenant
class PluginModelBackwardsInvocation(BaseBackwardsInvocation):
@staticmethod
@with_credit_usage_created_by(CreditUsageCreatedBy.PLUGIN_API)
def _get_bound_model_instance(
*,
tenant_id: str,
@@ -9,9 +9,11 @@ from sqlalchemy import select
from sqlalchemy.orm import Session
from core.app.file_access import DatabaseFileAccessController
from core.credit_usage import CreditUsageCreatedBy
from core.db.session_factory import session_factory
from core.entities.knowledge_entities import PreviewDetail
from core.llm_generator.prompts import DEFAULT_GENERATOR_SUMMARY_PROMPT
from core.model_context import with_credit_usage_created_by
from core.model_manager import ModelManager
from core.rag.cleaner.clean_processor import CleanProcessor
from core.rag.datasource.keyword.keyword_factory import Keyword
@@ -376,6 +378,7 @@ class ParagraphIndexProcessor(BaseIndexProcessor):
return preview_texts
@staticmethod
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_INDEXING)
def generate_summary(
tenant_id: str,
text: str,
+3
View File
@@ -3,6 +3,8 @@ from typing import override
from sqlalchemy.orm import Session
from core.credit_usage import CreditUsageCreatedBy
from core.model_context import with_credit_usage_created_by
from core.model_manager import ModelInstance, ModelManager
from core.rag.index_processor.constant.doc_type import DocType
from core.rag.index_processor.constant.query_type import QueryType
@@ -24,6 +26,7 @@ class RerankModelRunner(BaseRerankRunner):
@override
@trace_span()
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_RETRIEVAL)
def run(
self,
query: str,
+3
View File
@@ -4,6 +4,8 @@ from typing import override
import numpy as np
from core.credit_usage import CreditUsageCreatedBy
from core.model_context import with_credit_usage_created_by
from core.model_manager import ModelManager
from core.rag.datasource.keyword.jieba.jieba_keyword_table_handler import JiebaKeywordTableHandler
from core.rag.embedding.cached_embedding import CacheEmbedding
@@ -151,6 +153,7 @@ class WeightRerankRunner(BaseRerankRunner):
return similarities
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_RETRIEVAL)
def _calculate_cosine(
self, tenant_id: str, query: str, documents: list[Document], vector_setting: VectorSetting
) -> list[float]:
+24 -2
View File
@@ -19,13 +19,20 @@ from core.app.app_config.entities import (
MetadataFilteringCondition,
ModelConfig,
)
from core.app.entities.app_invoke_entities import InvokeFrom, ModelConfigWithCredentialsEntity
from core.app.entities.app_invoke_entities import (
CreditUsageCreatedBy,
EasyUIBasedAppGenerateEntity,
InvokeFrom,
ModelConfigWithCredentialsEntity,
get_credit_usage_app_type,
)
from core.app.file_access import grant_retriever_segment_access, grant_upload_file_access
from core.callback_handler.index_tool_callback_handler import DatasetIndexToolCallbackHandler
from core.db.session_factory import session_factory
from core.entities.agent_entities import PlanningStrategy
from core.entities.model_entities import ModelStatus
from core.memory.token_buffer_memory import TokenBufferMemory
from core.model_context import with_credit_usage_created_by, with_credit_usage_metadata
from core.model_manager import ModelInstance, ModelManager
from core.ops.entities.trace_entity import TraceTaskName
from core.ops.ops_trace_manager import TraceQueueManager, TraceTask
@@ -102,9 +109,19 @@ logger = logging.getLogger(__name__)
class DatasetRetrieval:
def __init__(self, application_generate_entity=None):
def __init__(self, application_generate_entity: EasyUIBasedAppGenerateEntity | None = None):
self.application_generate_entity = application_generate_entity
self._llm_usage = LLMUsage.empty_usage()
self._request_metadata: dict[str, object] | None = None
if application_generate_entity is not None:
app_config = application_generate_entity.app_config
self._request_metadata = {
"app_type": get_credit_usage_app_type(app_config.app_mode),
"app_id": app_config.app_id,
}
def set_request_metadata(self, request_metadata: Mapping[str, object] | None) -> None:
self._request_metadata = dict(request_metadata) if request_metadata else None
@property
def llm_usage(self) -> LLMUsage:
@@ -119,6 +136,8 @@ class DatasetRetrieval:
self._llm_usage = self._llm_usage.plus(usage)
@trace_span()
@with_credit_usage_metadata
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_RETRIEVAL)
def knowledge_retrieval(self, session: Session, request: KnowledgeRetrievalRequest) -> list[Source]:
self._check_knowledge_rate_limit(request.tenant_id)
available_datasets = self._get_available_datasets(request.tenant_id, request.dataset_ids)
@@ -354,6 +373,8 @@ class DatasetRetrieval:
item.metadata.position = position # type: ignore[index]
return retrieval_resource_list
@with_credit_usage_metadata
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_RETRIEVAL)
def retrieve(
self,
session: Session,
@@ -1447,6 +1468,7 @@ class DatasetRetrieval:
)
return filter_documents[:top_k] if top_k else filter_documents
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_RETRIEVAL)
def get_metadata_filter_condition(
self,
session: Session,
@@ -1,12 +1,15 @@
from typing import Union
from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity
from core.credit_usage import CreditUsageCreatedBy
from core.model_context import with_credit_usage_created_by
from core.model_manager import ModelInstance
from graphon.model_runtime.entities.llm_entities import LLMResult, LLMUsage
from graphon.model_runtime.entities.message_entities import PromptMessageTool, SystemPromptMessage, UserPromptMessage
class FunctionCallMultiDatasetRouter:
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_RETRIEVAL)
def invoke(
self,
query: str,
@@ -2,6 +2,8 @@ from collections.abc import Generator, Sequence
from typing import Any, Union
from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity
from core.credit_usage import CreditUsageCreatedBy
from core.model_context import with_credit_usage_created_by
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
@@ -135,6 +137,7 @@ class ReactMultiDatasetRouter:
return react_decision.tool, usage
return None, usage
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_RETRIEVAL)
def _invoke_llm(
self,
completion_param: dict[str, Any],
+5 -1
View File
@@ -6,6 +6,7 @@ from sqlalchemy import select
from core.db.session_factory import session_factory
from core.rag.index_processor.constant.index_type import IndexTechniqueType
from core.rag.index_processor.index_processor_base import SummaryIndexSettingDict
from extensions.otel import propagate_context
from models.dataset import Dataset, Document, DocumentSegment, DocumentSegmentSummary
from services.summary_index_service import SummaryIndexService
from tasks.generate_summary_index_task import generate_summary_index_task
@@ -92,7 +93,10 @@ class SummaryIndex:
# Continue processing other segments
with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor:
futures = [executor.submit(process_segment, segment_id) for segment_id in pending_segment_ids]
futures = [
executor.submit(propagate_context(process_segment), segment_id)
for segment_id in pending_segment_ids
]
concurrent.futures.wait(futures)
else:
generate_summary_index_task.delay(dataset_id, document_id, None)
@@ -4,6 +4,8 @@ from typing import Any, override
from sqlalchemy.orm import Session
from core.credit_usage import CreditUsageCreatedBy
from core.model_context import with_credit_usage_created_by
from core.model_manager import ModelManager
from core.plugin.entities.parameters import PluginParameterOption
from core.tools.builtin_tool.tool import BuiltinTool
@@ -17,6 +19,7 @@ from services.model_provider_service import ModelProviderService
class ASRTool(BuiltinTool):
@override
@with_credit_usage_created_by(CreditUsageCreatedBy.AUDIO)
def _invoke(
self,
session: Session,
@@ -4,6 +4,8 @@ from typing import Any, override
from sqlalchemy.orm import Session
from core.base.tts.audio_mime import get_model_audio_mime_type, inspect_audio_stream
from core.credit_usage import CreditUsageCreatedBy
from core.model_context import with_credit_usage_created_by
from core.model_manager import ModelManager
from core.plugin.entities.parameters import PluginParameterOption
from core.tools.builtin_tool.tool import BuiltinTool
@@ -15,6 +17,7 @@ from services.model_provider_service import ModelProviderService
class TTSTool(BuiltinTool):
@override
@with_credit_usage_created_by(CreditUsageCreatedBy.AUDIO)
def _invoke(
self,
session: Session,
@@ -8,6 +8,8 @@ import json
from decimal import Decimal
from typing import cast
from core.credit_usage import CreditUsageCreatedBy
from core.model_context import with_credit_usage_created_by
from core.model_manager import ModelManager
from core.tools.entities.tool_entities import ToolProviderType
from extensions.ext_database import db
@@ -79,6 +81,7 @@ class ModelInvocationUtils:
return tokens
@staticmethod
@with_credit_usage_created_by(CreditUsageCreatedBy.TOOL)
def invoke(
user_id: str,
tenant_id: str,
+11 -1
View File
@@ -10,6 +10,7 @@ from sqlalchemy import select
from configs import dify_config
from core.app.entities.app_invoke_entities import DIFY_RUN_CONTEXT_KEY, DifyRunContext
from core.app.llm.model_access import build_dify_model_access, fetch_model_config
from core.credit_usage import created_by_from_app_type
from core.db.session_factory import session_factory
from core.file import remote_fetcher
from core.helper.code_executor.code_executor import (
@@ -616,11 +617,20 @@ 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)
request_metadata: dict[str, object] = {"app_id": self._dify_context.app_id}
app_type = self._dify_context.app_type
created_by = self._dify_context.created_by
if app_type is not None:
request_metadata["app_type"] = app_type
request_metadata["created_by"] = created_by_from_app_type(app_type)
elif created_by is not None:
request_metadata["created_by"] = created_by
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},
request_metadata=request_metadata,
)
if wrap_model_instance
else model_instance
+26 -20
View File
@@ -20,6 +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_context import use_credit_usage_metadata
from core.model_manager import ModelInstance, QuotaManagedModelInstance
from core.plugin.impl.exc import PluginDaemonClientSideError, PluginInvokeError
from core.plugin.impl.plugin import PluginInstaller
@@ -303,6 +304,7 @@ class DifyPreparedLLM(LLMProtocol):
model_parameters=model_parameters,
stop=list(stop or []),
stream=stream,
request_metadata=self._request_metadata,
)
@override
@@ -362,7 +364,7 @@ class DifyPreparedPollingLLM(DifyPreparedLLM, LLMPollingCapableProtocol):
self.finalize_llm_polling()
if isinstance(self._model_instance, QuotaManagedModelInstance):
self._polling_quota_reservation = self._model_instance.reserve_quota()
self._polling_quota_reservation = self._model_instance._reserve_quota_for_request(self._request_metadata)
try:
polling_result = self._polling_runtime.start_llm_polling(
@@ -627,25 +629,29 @@ class DifyToolNodeRuntime(ToolNodeRuntimeProtocol):
tool.clear_trace_session_id()
try:
session_maker = self._session_maker or session_factory.get_session_maker()
with session_maker.begin() as session:
messages = ToolEngine.generic_invoke(
session=session,
tool=tool,
tool_parameters=dict(tool_parameters),
user_id=self._run_context.user_id,
workflow_tool_callback=callback,
workflow_call_depth=workflow_call_depth,
app_id=self._run_context.app_id,
conversation_id=runtime_binding.conversation_id,
)
transformed_messages = ToolFileMessageTransformer.transform_tool_invoke_messages(
messages=messages,
user_id=self._run_context.user_id,
tenant_id=self._run_context.tenant_id,
conversation_id=runtime_binding.conversation_id,
)
yield from self._adapt_messages(transformed_messages, provider_name=provider_name)
request_metadata = (
{"app_type": self._run_context.app_type} if self._run_context.app_type is not None else None
)
with use_credit_usage_metadata(request_metadata):
session_maker = self._session_maker or session_factory.get_session_maker()
with session_maker.begin() as session:
messages = ToolEngine.generic_invoke(
session=session,
tool=tool,
tool_parameters=dict(tool_parameters),
user_id=self._run_context.user_id,
workflow_tool_callback=callback,
workflow_call_depth=workflow_call_depth,
app_id=self._run_context.app_id,
conversation_id=runtime_binding.conversation_id,
)
transformed_messages = ToolFileMessageTransformer.transform_tool_invoke_messages(
messages=messages,
user_id=self._run_context.user_id,
tenant_id=self._run_context.tenant_id,
conversation_id=runtime_binding.conversation_id,
)
yield from self._adapt_messages(transformed_messages, provider_name=provider_name)
except Exception as exc:
raise self._map_invocation_exception(exc, provider_name=provider_name) from exc
@@ -4,7 +4,9 @@ from typing import TYPE_CHECKING, Any, override
from sqlalchemy.orm import Session
from core.app.entities.app_invoke_entities import DIFY_RUN_CONTEXT_KEY, DifyRunContext
from core.db.session_factory import session_factory
from core.model_context import use_credit_usage_metadata
from core.rag.index_processor.index_processor import IndexProcessor
from core.rag.index_processor.index_processor_base import SummaryIndexSettingDict
from core.rag.summary_index.summary_index import SummaryIndex
@@ -150,14 +152,17 @@ class KnowledgeIndexNode(Node[KnowledgeIndexNodeData]):
):
if not document_id:
raise KnowledgeIndexNodeError("document_id is required.")
rst = self.index_processor.index_and_clean(
dataset_id, document_id, original_document_id, chunks, batch, summary_index_setting, session=session
)
# Summary generation opens independent sessions and must see the indexed rows.
session.commit()
self.summary_index_service.generate_and_vectorize_summary(
dataset_id, document_id, is_preview, summary_index_setting
)
dify_ctx = DifyRunContext.model_validate(self.require_run_context_value(DIFY_RUN_CONTEXT_KEY))
request_metadata = {"app_type": dify_ctx.app_type} if dify_ctx.app_type is not None else None
with use_credit_usage_metadata(request_metadata):
rst = self.index_processor.index_and_clean(
dataset_id, document_id, original_document_id, chunks, batch, summary_index_setting, session=session
)
# Summary generation opens independent sessions and must see the indexed rows.
session.commit()
self.summary_index_service.generate_and_vectorize_summary(
dataset_id, document_id, is_preview, summary_index_setting
)
return rst
@classmethod
@@ -12,6 +12,7 @@ from sqlalchemy.orm import Session
from core.app.app_config.entities import DatasetRetrieveConfigEntity
from core.app.entities.app_invoke_entities import DIFY_RUN_CONTEXT_KEY, DifyRunContext
from core.credit_usage import CreditUsageAppType
from core.db.session_factory import session_factory
from core.rag.data_post_processor.data_post_processor import RerankingModelDict, WeightsDict
from core.rag.retrieval.dataset_retrieval import DatasetRetrieval
@@ -185,6 +186,12 @@ class KnowledgeRetrievalNode(Node[KnowledgeRetrievalNodeData]):
self, session: Session, node_data: KnowledgeRetrievalNodeData, variables: dict[str, Any]
) -> tuple[list[Source], LLMUsage]:
dify_ctx = DifyRunContext.model_validate(self.require_run_context_value(DIFY_RUN_CONTEXT_KEY))
self._rag_retrieval.set_request_metadata(
{
"app_id": dify_ctx.app_id,
"app_type": dify_ctx.app_type or CreditUsageAppType.UNKNOWN,
}
)
dataset_ids = node_data.dataset_ids
query = variables.get("query")
attachments = variables.get("attachments")
@@ -1,3 +1,4 @@
from collections.abc import Mapping
from typing import Any, Literal, Protocol
from pydantic import BaseModel, Field
@@ -86,4 +87,6 @@ class RAGRetrievalProtocol(Protocol):
@property
def llm_usage(self) -> LLMUsage: ...
def set_request_metadata(self, request_metadata: Mapping[str, object] | None) -> None: ...
def knowledge_retrieval(self, request: KnowledgeRetrievalRequest) -> list[Source]: ...
+8 -1
View File
@@ -7,9 +7,14 @@ from uuid import uuid4
from configs import dify_config
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.entities.app_invoke_entities import (
InvokeFrom,
UserFrom,
build_dify_run_context,
)
from core.app.file_access import DatabaseFileAccessController
from core.app.workflow.layers.observability import ObservabilityLayer
from core.credit_usage import CreditUsageAppType
from core.workflow.node_factory import (
DifyGraphInitContext,
DifyNodeFactory,
@@ -228,6 +233,7 @@ class WorkflowEntry:
user_id=user_id,
user_from=UserFrom.ACCOUNT,
invoke_from=InvokeFrom.DEBUGGER,
app_type=CreditUsageAppType.WORKFLOW,
)
graph_init_context = DifyGraphInitContext(
workflow_id=workflow.id,
@@ -387,6 +393,7 @@ class WorkflowEntry:
user_id=user_id,
user_from=UserFrom.ACCOUNT,
invoke_from=InvokeFrom.DEBUGGER,
app_type=CreditUsageAppType.WORKFLOW,
)
graph_init_context = DifyGraphInitContext(
workflow_id="",
@@ -1,7 +1,8 @@
import logging
import time as time_module
from collections.abc import Mapping
from datetime import datetime
from typing import Any, cast
from typing import Any, TypedDict, cast
from pydantic import BaseModel
from sqlalchemy import update
@@ -13,11 +14,14 @@ from core.app.entities.app_invoke_entities import (
AgentAppGenerateEntity,
AgentChatAppGenerateEntity,
ChatAppGenerateEntity,
get_credit_usage_app_type,
get_credit_usage_created_by,
)
from core.entities.provider_entities import ProviderQuotaType, QuotaUnit, SystemConfiguration
from events.message_event import message_was_created
from extensions.ext_database import db
from extensions.ext_redis import redis_client, redis_fallback
from graphon.model_runtime.entities.model_entities import ModelType
from libs import datetime_utils
from models.model import Message
from models.provider import Provider, ProviderType
@@ -83,6 +87,11 @@ class _ProviderUpdateOperation(BaseModel):
description: str = "unknown"
class _CreditDeductionContext(TypedDict, total=False):
request_id: str | None
metadata: Mapping[str, object]
@message_was_created.connect
def handle(sender: Message, **kwargs):
"""
@@ -124,7 +133,18 @@ def handle(sender: Message, **kwargs):
model_config = application_generate_entity.model_conf
provider_model_bundle = model_config.provider_model_bundle
provider_configuration = provider_model_bundle.configuration
app_mode = application_generate_entity.app_config.app_mode
credit_deduction_metadata: dict[str, object] = {
"provider": provider_name,
"model": model_config.model,
"model_type": ModelType.LLM.value,
"app_type": get_credit_usage_app_type(app_mode),
"created_by": get_credit_usage_created_by(app_mode),
}
credit_deduction_context: _CreditDeductionContext = {
"request_id": str(message.id) if message.id else None,
"metadata": credit_deduction_metadata,
}
agent_gateway_metered = (
isinstance(application_generate_entity, AgentAppGenerateEntity)
and application_generate_entity.agent_llm_gateway_enabled
@@ -150,12 +170,14 @@ def handle(sender: Message, **kwargs):
tenant_id=tenant_id,
credits_required=used_quota,
pool_type="trial",
**credit_deduction_context,
)
case ProviderQuotaType.PAID:
_deduct_credit_pool_quota_capped(
tenant_id=tenant_id,
credits_required=used_quota,
pool_type="paid",
**credit_deduction_context,
)
case ProviderQuotaType.FREE:
quota_update = _ProviderUpdateOperation(
@@ -205,16 +227,33 @@ def handle(sender: Message, **kwargs):
raise
def _deduct_credit_pool_quota_capped(*, tenant_id: str, credits_required: int, pool_type: str) -> None:
def _deduct_credit_pool_quota_capped(
*,
tenant_id: str,
credits_required: int,
pool_type: str,
request_id: str | None = None,
metadata: Mapping[str, object] | None = None,
) -> None:
"""Apply post-generation credit accounting without failing message persistence on quota exhaustion."""
from services.credit_pool_service import CreditPoolService
deducted_credits = CreditPoolService.deduct_credits_capped(
tenant_id=tenant_id,
credits_required=credits_required,
pool_type=pool_type,
session=db.session(),
)
if request_id is None and metadata is None:
deducted_credits = CreditPoolService.deduct_credits_capped(
tenant_id=tenant_id,
credits_required=credits_required,
pool_type=pool_type,
session=db.session(),
)
else:
deducted_credits = CreditPoolService.deduct_credits_capped(
tenant_id=tenant_id,
credits_required=credits_required,
pool_type=pool_type,
request_id=request_id,
metadata=metadata,
session=db.session(),
)
if deducted_credits < credits_required:
logger.warning(
"Credit pool exhausted during message-created accounting, "
+3 -1
View File
@@ -1,5 +1,6 @@
"""Utilities for propagating OpenTelemetry context across execution boundaries."""
import contextvars
import functools
from collections.abc import Callable
@@ -9,12 +10,13 @@ from opentelemetry import context as otel_context
def propagate_context[**P, R](func: Callable[P, R]) -> Callable[P, R]:
"""Capture the current context and attach it whenever ``func`` executes."""
captured_context = otel_context.get_current()
captured_contextvars = contextvars.copy_context()
@functools.wraps(func)
def wrapper(*args: P.args, **kwargs: P.kwargs) -> R:
token = otel_context.attach(captured_context)
try:
return func(*args, **kwargs)
return captured_contextvars.run(func, *args, **kwargs)
finally:
otel_context.detach(token)
+22 -3
View File
@@ -8,6 +8,8 @@ from typing import cast
from sqlalchemy.orm import Session
from core.app.entities.app_invoke_entities import get_credit_usage_app_type
from core.credit_usage import CreditUsageAppType, CreditUsageCreatedBy
from core.db.session_factory import session_factory as default_session_factory
from core.entities.model_entities import ModelStatus
from core.model_manager import ModelInstance, ModelManager
@@ -31,6 +33,15 @@ class AgentLLMInnerServiceError(RuntimeError):
class PreparedAgentLLMInvocation:
request: AgentLLMInvokeRequest
model_instance: ModelInstance
app_type: CreditUsageAppType = CreditUsageAppType.UNKNOWN
@property
def created_by(self) -> CreditUsageCreatedBy:
if self.request.caller.agent_config_version_kind == "build_draft":
return CreditUsageCreatedBy.BUILD_DRAFT
if self.app_type is CreditUsageAppType.AGENT_V2:
return CreditUsageCreatedBy.APP
return CreditUsageCreatedBy.AGENT_NODE
class AgentLLMInnerService:
@@ -42,7 +53,7 @@ class AgentLLMInnerService:
def prepare(self, request: AgentLLMInvokeRequest) -> PreparedAgentLLMInvocation:
caller = request.caller
target = request.target
self._validate_app_tenant(app_id=caller.app_id, tenant_id=caller.tenant_id)
app = self._validate_app_tenant(app_id=caller.app_id, tenant_id=caller.tenant_id)
provider_manager = create_plugin_provider_manager(tenant_id=caller.tenant_id, user_id=caller.user_id)
model_manager = ModelManager(provider_manager=provider_manager)
model_instance = model_manager.get_model_instance(
@@ -65,7 +76,11 @@ class AgentLLMInnerService:
if provider_model.status != ModelStatus.QUOTA_EXCEEDED:
provider_model.raise_for_status()
return PreparedAgentLLMInvocation(request=request, model_instance=model_instance)
return PreparedAgentLLMInvocation(
request=request,
model_instance=model_instance,
app_type=get_credit_usage_app_type(app.mode),
)
def invoke(self, prepared: PreparedAgentLLMInvocation) -> Generator[LLMResultChunk, None, None]:
request = prepared.request
@@ -84,16 +99,19 @@ class AgentLLMInnerService:
"invocation_id": caller.invocation_id,
"agent_run_id": caller.agent_run_id,
"agent_mode": caller.agent_mode,
"agent_config_version_kind": caller.agent_config_version_kind,
"call_index": caller.call_index,
"app_id": caller.app_id,
"workflow_run_id": caller.workflow_run_id,
"node_execution_id": caller.node_execution_id,
"trace_id": caller.trace_id,
"app_type": prepared.app_type,
"created_by": prepared.created_by,
},
)
yield from cast(Generator[LLMResultChunk, None, None], result)
def _validate_app_tenant(self, *, app_id: str, tenant_id: str) -> None:
def _validate_app_tenant(self, *, app_id: str, tenant_id: str) -> App:
with self._session_factory() as session:
app = session.get(App, app_id)
if app is None:
@@ -108,6 +126,7 @@ class AgentLLMInnerService:
"App does not belong to the caller tenant.",
status_code=403,
)
return app
__all__ = ["AgentLLMInnerService", "AgentLLMInnerServiceError", "PreparedAgentLLMInvocation"]
+22 -3
View File
@@ -11,7 +11,9 @@ from werkzeug.datastructures import FileStorage
from constants import AUDIO_EXTENSIONS
from core.app.apps.agent_app.app_feature_projection import merge_agent_app_features
from core.app.entities.app_invoke_entities import get_credit_usage_app_type
from core.base.tts.audio_mime import get_model_audio_mime_type, inspect_audio_stream, resolve_audio_mime_type
from core.credit_usage import CreditUsageCreatedBy
from core.model_manager import ModelManager
from graphon.model_runtime.entities.model_entities import ModelType
from models.agent_config_entities import AgentSoulConfig
@@ -160,7 +162,14 @@ class AudioService:
message = f"Audio size larger than {FILE_SIZE} mb"
raise AudioTooLargeServiceError(message)
model_manager = ModelManager.for_tenant(tenant_id=app_model.tenant_id, user_id=end_user)
model_manager = ModelManager.for_tenant(
tenant_id=app_model.tenant_id,
user_id=end_user,
request_metadata={
"app_type": get_credit_usage_app_type(app_model.mode),
"created_by": CreditUsageCreatedBy.AUDIO,
},
)
model_instance = model_manager.get_default_model_instance(
tenant_id=app_model.tenant_id, model_type=ModelType.SPEECH2TEXT
)
@@ -211,7 +220,14 @@ class AudioService:
voice = cast(str | None, text_to_speech_dict.get("voice"))
model_manager = ModelManager.for_tenant(tenant_id=app_model.tenant_id, user_id=end_user)
model_manager = ModelManager.for_tenant(
tenant_id=app_model.tenant_id,
user_id=end_user,
request_metadata={
"app_type": get_credit_usage_app_type(app_model.mode),
"created_by": CreditUsageCreatedBy.AUDIO,
},
)
model_instance = model_manager.get_default_model_instance(
tenant_id=app_model.tenant_id, model_type=ModelType.TTS
)
@@ -258,7 +274,10 @@ class AudioService:
@classmethod
def transcript_tts_voices(cls, tenant_id: str, language: str):
model_manager = ModelManager.for_tenant(tenant_id=tenant_id)
model_manager = ModelManager.for_tenant(
tenant_id=tenant_id,
request_metadata={"created_by": CreditUsageCreatedBy.AUDIO},
)
model_instance = model_manager.get_default_model_instance(tenant_id=tenant_id, model_type=ModelType.TTS)
if model_instance is None:
raise ProviderNotSupportTextToSpeechServiceError()
+6 -2
View File
@@ -7,8 +7,9 @@ from sqlalchemy import asc, desc, func, or_, select
from sqlalchemy.orm import Session
from configs import dify_config
from core.app.entities.app_invoke_entities import InvokeFrom
from core.app.entities.app_invoke_entities import InvokeFrom, get_credit_usage_app_type
from core.llm_generator.llm_generator import LLMGenerator
from core.model_context import use_credit_usage_metadata
from factories import variable_factory
from graphon.variables.types import SegmentType
from libs.datetime_utils import naive_utc_now
@@ -152,7 +153,10 @@ class ConversationService:
raise MessageNotExistsError()
# generate conversation name
with contextlib.suppress(Exception):
with (
contextlib.suppress(Exception),
use_credit_usage_metadata({"app_type": get_credit_usage_app_type(app_model.mode)}),
):
name = LLMGenerator.generate_conversation_name(
app_model.tenant_id, message.query, conversation.id, app_model.id
)
+28 -5
View File
@@ -16,6 +16,7 @@ from sqlalchemy import select
from sqlalchemy.orm import Session
from configs import dify_config
from core.credit_usage import normalize_credit_usage_app_type, normalize_credit_usage_created_by
from core.errors.error import QuotaExceededError
from enums import DeploymentEdition
from extensions.ext_redis import redis_client
@@ -25,6 +26,8 @@ from models.enums import ProviderQuotaType
logger = logging.getLogger(__name__)
FEATURE_KEY_CREDIT_POOL = "credit_pool"
CREDIT_USAGE_CREATED_BY_META_KEY = "created_by"
CREDIT_USAGE_APP_TYPE_META_KEY = "app_type"
CREDIT_POOL_TENANT_LOCK_TIMEOUT_SECONDS = 10
CREDIT_POOL_TENANT_LOCK_BLOCKING_TIMEOUT_SECONDS = 5
@@ -123,6 +126,23 @@ class CreditPoolService:
def _normalize_pool_type(pool_type: str | ProviderQuotaType) -> str:
return pool_type.value if isinstance(pool_type, ProviderQuotaType) else str(pool_type)
@staticmethod
def _build_billing_metadata(
source: str,
metadata: Mapping[str, object] | None,
) -> dict[str, object]:
billing_metadata: dict[str, object] = {
"source": source,
**dict(metadata or {}),
}
billing_metadata[CREDIT_USAGE_CREATED_BY_META_KEY] = normalize_credit_usage_created_by(
billing_metadata.get(CREDIT_USAGE_CREATED_BY_META_KEY)
).value
billing_metadata[CREDIT_USAGE_APP_TYPE_META_KEY] = normalize_credit_usage_app_type(
billing_metadata.get(CREDIT_USAGE_APP_TYPE_META_KEY)
).value
return billing_metadata
@staticmethod
def _use_billing_quota() -> bool:
return bool(dify_config.DEPLOYMENT_EDITION == DeploymentEdition.CLOUD)
@@ -254,7 +274,7 @@ class CreditPoolService:
raise ValueError("request_id is required")
normalized_pool_type = cls._normalize_pool_type(pool_type)
reservation_meta = {"source": "credit_pool.reservation", **(meta or {})}
reservation_meta = cls._build_billing_metadata("credit_pool.reservation", meta)
if cls._use_billing_quota():
from services.billing_service import BillingService
@@ -348,7 +368,7 @@ class CreditPoolService:
pool_type: str | ProviderQuotaType = "trial",
*,
request_id: str | None = None,
metadata: Mapping[str, str] | None = None,
metadata: Mapping[str, object] | None = None,
session: Session | None = None,
) -> int:
"""Deduct exactly the requested credits or raise without mutating the pool."""
@@ -360,7 +380,7 @@ class CreditPoolService:
from services.billing_service import BillingService
resolved_request_id = request_id or str(uuid4())
billing_metadata = {"source": "credit_pool.check_and_deduct", **dict(metadata or {})}
billing_metadata = cls._build_billing_metadata("credit_pool.check_and_deduct", metadata)
result = BillingService.quota_reserve(
tenant_id=tenant_id,
feature_key=FEATURE_KEY_CREDIT_POOL,
@@ -432,6 +452,8 @@ class CreditPoolService:
credits_required: int,
pool_type: str | ProviderQuotaType = "trial",
*,
request_id: str | None = None,
metadata: Mapping[str, object] | None = None,
session: Session | None = None,
) -> int:
"""Deduct up to the available balance and return the actual deducted credits."""
@@ -442,13 +464,14 @@ class CreditPoolService:
if cls._use_billing_quota():
from services.billing_service import BillingService
billing_metadata = cls._build_billing_metadata("credit_pool.deduct_capped", metadata)
result = BillingService.quota_consume_capped(
tenant_id=tenant_id,
feature_key=FEATURE_KEY_CREDIT_POOL,
bucket=normalized_pool_type,
request_id=str(uuid4()),
request_id=request_id or str(uuid4()),
amount=credits_required,
meta={"source": "credit_pool.deduct_capped"},
meta=billing_metadata,
)
return result["deducted"]
+6 -2
View File
@@ -7,9 +7,10 @@ from sqlalchemy.orm import Session, sessionmaker
from core.app.apps.advanced_chat.app_config_manager import AdvancedChatAppConfigManager
from core.app.apps.agent_app.app_feature_projection import merge_agent_app_features
from core.app.entities.app_invoke_entities import InvokeFrom
from core.app.entities.app_invoke_entities import InvokeFrom, get_credit_usage_app_type
from core.llm_generator.llm_generator import LLMGenerator
from core.memory.token_buffer_memory import TokenBufferMemory
from core.model_context import use_credit_usage_metadata
from core.model_manager import ModelManager
from core.ops.entities.trace_entity import TraceTaskName
from core.ops.ops_trace_manager import TraceQueueManager, TraceTask
@@ -428,7 +429,10 @@ class MessageService:
instruction_prompt = None
configured_model = suggested_questions_after_answer_config.get("model")
with measure_time() as timer:
with (
measure_time() as timer,
use_credit_usage_metadata({"app_type": get_credit_usage_app_type(app_model.mode)}),
):
questions_sequence = LLMGenerator.generate_suggested_questions_after_answer(
tenant_id=app_model.tenant_id,
histories=histories,
+5
View File
@@ -9,7 +9,9 @@ from typing import TypedDict, cast
from sqlalchemy import select, update
from sqlalchemy.orm import Session
from core.credit_usage import CreditUsageCreatedBy
from core.db.session_factory import session_factory
from core.model_context import with_credit_usage_created_by
from core.model_manager import ModelManager
from core.rag.datasource.vdb.vector_factory import Vector
from core.rag.index_processor.constant.doc_type import DocType
@@ -46,6 +48,7 @@ class SummaryIndexService:
"""Service for generating and managing summary indexes."""
@staticmethod
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_INDEXING)
def generate_summary_for_segment(
segment: DocumentSegment,
dataset: Dataset,
@@ -158,6 +161,7 @@ class SummaryIndexService:
return summary_record
@staticmethod
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_INDEXING)
def vectorize_summary(
summary_record: DocumentSegmentSummary,
segment: DocumentSegment,
@@ -655,6 +659,7 @@ class SummaryIndexService:
logger.warning("Summary record not found for segment %s when updating error", segment.id)
@staticmethod
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_INDEXING)
def generate_and_vectorize_summary(
segment: DocumentSegment,
dataset: Dataset,
+8
View File
@@ -3,6 +3,8 @@ import logging
from sqlalchemy import delete, select
from sqlalchemy.orm import Session
from core.credit_usage import CreditUsageCreatedBy
from core.model_context import with_credit_usage_created_by
from core.model_manager import ModelInstance, ModelManager
from core.rag.datasource.keyword.keyword_factory import Keyword
from core.rag.datasource.vdb.vector_factory import Vector
@@ -23,6 +25,7 @@ logger = logging.getLogger(__name__)
class VectorService:
@classmethod
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_INDEXING)
def create_segments_vector(
cls,
keywords_list: list[list[str]] | None,
@@ -112,6 +115,7 @@ class VectorService:
index_processor.load(dataset, [], multimodal_documents, with_keywords=False, session=session)
@classmethod
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_INDEXING)
def update_segment_vector(
cls, keywords: list[str] | None, segment: DocumentSegment, dataset: Dataset, session: Session
):
@@ -145,6 +149,7 @@ class VectorService:
keyword.add_texts([document], session)
@classmethod
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_INDEXING)
def generate_child_chunks(
cls,
segment: DocumentSegment,
@@ -214,6 +219,7 @@ class VectorService:
session.flush()
@classmethod
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_INDEXING)
def create_child_chunk_vector(cls, child_segment: ChildChunk, dataset: Dataset, *, session: Session):
child_document = Document(
page_content=child_segment.content,
@@ -230,6 +236,7 @@ class VectorService:
vector.add_texts([child_document], duplicate_check=True)
@classmethod
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_INDEXING)
def update_child_chunk_vector(
cls,
new_child_chunks: list[ChildChunk],
@@ -283,6 +290,7 @@ class VectorService:
vector.delete_by_ids([child_chunk.index_node_id])
@classmethod
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_INDEXING)
def update_multimodel_vector(
cls, segment: DocumentSegment, attachment_ids: list[str], dataset: Dataset, session: Session
):
+15 -3
View File
@@ -16,6 +16,7 @@ from collections.abc import Iterator
from typing import Any
from core.app.app_config.entities import ModelConfig
from core.credit_usage import CreditUsageAppType, CreditUsageCreatedBy
from core.model_manager import ModelInstance, ModelManager
from core.workflow.generator import WorkflowGenerator
from core.workflow.generator.tool_catalogue import (
@@ -67,7 +68,7 @@ class WorkflowGeneratorService:
envelope as ``/rule-generate``).
"""
model_instance, model_parameters, tool_catalogue_entries, tool_catalogue_text, installed_tools = (
cls._resolve_generation_context(tenant_id=tenant_id, model_config=model_config)
cls._resolve_generation_context(tenant_id=tenant_id, mode=mode, model_config=model_config)
)
return WorkflowGenerator.generate_workflow_graph(
@@ -107,7 +108,7 @@ class WorkflowGeneratorService:
single ``result`` SSE event).
"""
model_instance, model_parameters, tool_catalogue_entries, tool_catalogue_text, installed_tools = (
cls._resolve_generation_context(tenant_id=tenant_id, model_config=model_config)
cls._resolve_generation_context(tenant_id=tenant_id, mode=mode, model_config=model_config)
)
yield from WorkflowGenerator.generate_workflow_graph_stream(
@@ -130,6 +131,7 @@ class WorkflowGeneratorService:
cls,
*,
tenant_id: str,
mode: WorkflowGenerationModeRequest,
model_config: ModelConfig,
) -> tuple[
ModelInstance,
@@ -149,7 +151,17 @@ class WorkflowGeneratorService:
empty set, so we don't reject every tool node just because we couldn't
enumerate the catalogue).
"""
model_manager = ModelManager.for_tenant(tenant_id=tenant_id)
app_type = {
"workflow": CreditUsageAppType.WORKFLOW,
"advanced-chat": CreditUsageAppType.CHATFLOW,
}.get(mode, CreditUsageAppType.UNKNOWN)
model_manager = ModelManager.for_tenant(
tenant_id=tenant_id,
request_metadata={
"app_type": app_type,
"created_by": CreditUsageCreatedBy.WORKFLOW_GENERATION,
},
)
model_instance = model_manager.get_model_instance(
tenant_id=tenant_id,
model_type=ModelType.LLM,
@@ -10,7 +10,9 @@ import pandas as pd
from celery import shared_task
from sqlalchemy import func, select
from core.credit_usage import CreditUsageCreatedBy
from core.db.session_factory import session_factory
from core.model_context import with_credit_usage_created_by
from core.model_manager import ModelManager
from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType
from extensions.ext_redis import redis_client
@@ -27,6 +29,7 @@ logger = logging.getLogger(__name__)
@shared_task(queue="dataset")
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_INDEXING)
def batch_create_segment_to_index_task(
job_id: str,
upload_file_id: str,
@@ -8,7 +8,9 @@ import click
from celery import shared_task
from sqlalchemy import or_, select
from core.credit_usage import CreditUsageCreatedBy
from core.db.session_factory import session_factory
from core.model_context import with_credit_usage_created_by
from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType
from models.dataset import Dataset, DocumentSegment, DocumentSegmentSummary
from models.dataset import Document as DatasetDocument
@@ -19,6 +21,7 @@ logger = logging.getLogger(__name__)
@shared_task(queue="dataset_summary")
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_INDEXING)
def regenerate_summary_index_task(
dataset_id: str,
regenerate_reason: str = "summary_model_changed",
@@ -12,6 +12,8 @@ from sqlalchemy.orm import Session
from core.agent.cot_agent_runner import CotAgentRunner
from core.agent.entities import AgentScratchpadUnit
from core.agent.errors import AgentMaxIterationError
from core.app.entities.app_invoke_entities import CreditUsageCreatedBy
from core.credit_usage import CreditUsageAppType
from graphon.model_runtime.entities.llm_entities import LLMUsage
from libs.datetime_utils import naive_utc_now
from models.enums import ConversationFromSource, MessageStatus
@@ -404,7 +406,11 @@ class TestRun:
results = list(runner.run(session, message, "query", {}))
assert events == ["commit", "close", "first-chunk"]
assert runner.model_instance.invoke_llm.call_args.kwargs["request_metadata"] == {"app_id": "app"}
assert runner.model_instance.invoke_llm.call_args.kwargs["request_metadata"] == {
"app_id": "app",
"app_type": CreditUsageAppType.AGENT,
"created_by": CreditUsageCreatedBy.APP.value,
}
assert results[-1].delta.message.content == ""
def test_run_usage_missing_key_branch(self, runner: DummyRunner, mocker: MockerFixture):
@@ -14,7 +14,9 @@ from sqlalchemy.orm import Session
from core.agent.errors import AgentMaxIterationError
from core.agent.fc_agent_runner import FunctionCallAgentRunner
from core.app.apps.base_app_queue_manager import PublishFrom
from core.app.entities.app_invoke_entities import CreditUsageCreatedBy
from core.app.entities.queue_entities import QueueMessageFileEvent
from core.credit_usage import CreditUsageAppType
from graphon.model_runtime.entities.llm_entities import LLMUsage
from graphon.model_runtime.entities.message_entities import (
DocumentPromptMessageContent,
@@ -435,7 +437,11 @@ class TestRunMethod:
assert len(outputs) == 1
assert "session" not in runner.create_agent_thought.call_args.kwargs
assert "session" not in runner.save_agent_thought.call_args.kwargs
assert runner.model_instance.invoke_llm.call_args.kwargs["request_metadata"] == {"app_id": "app"}
assert runner.model_instance.invoke_llm.call_args.kwargs["request_metadata"] == {
"app_id": "app",
"app_type": CreditUsageAppType.AGENT,
"created_by": CreditUsageCreatedBy.APP.value,
}
runner.queue_manager.publish.assert_called()
queue_calls = runner.queue_manager.publish.call_args_list
@@ -269,6 +269,7 @@ class TestChatAppRunner:
app_config = SimpleNamespace(
app_id="app-1",
tenant_id="tenant-1",
app_mode=AppMode.CHAT,
prompt_template=None,
external_data_variables=[],
dataset=None,
@@ -311,6 +312,7 @@ class TestChatAppRunner:
app_config = SimpleNamespace(
app_id="app-1",
tenant_id="tenant-1",
app_mode=AppMode.CHAT,
prompt_template=None,
external_data_variables=[],
dataset=None,
@@ -394,6 +396,7 @@ class TestChatAppRunner:
app_config = SimpleNamespace(
app_id="app-1",
tenant_id="tenant-1",
app_mode=AppMode.CHAT,
prompt_template=None,
external_data_variables=[],
dataset=None,
@@ -7,6 +7,7 @@ from sqlalchemy.orm import Session
import core.app.apps.completion.app_runner as module
from core.app.apps.completion.app_runner import CompletionAppRunner
from core.credit_usage import CreditUsageAppType, CreditUsageCreatedBy
from core.moderation.base import ModerationError
from graphon.model_runtime.entities.message_entities import ImagePromptMessageContent
from graphon.model_runtime.entities.model_entities import ModelType
@@ -25,6 +26,7 @@ def _build_app_config(dataset=None, external_tools=None, additional_features=Non
app_config = MagicMock()
app_config.app_id = APP_ID
app_config.tenant_id = TENANT_ID
app_config.app_mode = AppMode.COMPLETION
app_config.prompt_template = MagicMock()
app_config.dataset = dataset
app_config.external_data_variables = external_tools or []
@@ -245,7 +247,11 @@ class TestCompletionAppRunner:
model_parameters={"max_tokens": 10},
stop=["stop"],
stream=stream,
request_metadata={"app_id": APP_ID},
request_metadata={
"app_id": APP_ID,
"app_type": CreditUsageAppType.COMPLETION,
"created_by": CreditUsageCreatedBy.APP,
},
)
def test_run_uses_low_image_detail_default(self, runner, mocker: MockerFixture, sqlite_session: Session):
@@ -10,9 +10,11 @@ from core.app.apps.base_app_queue_manager import AppQueueManager
from core.app.apps.workflow.app_runner import WorkflowAppRunner
from core.app.apps.workflow_app_runner import WorkflowBasedAppRunner
from core.app.entities.app_invoke_entities import InvokeFrom, WorkflowAppGenerateEntity
from core.credit_usage import CreditUsageAppType
from core.workflow.system_variables import default_system_variables
from graphon.entities.graph_config import NodeConfigDictAdapter
from graphon.runtime import GraphRuntimeState, VariablePool
from models.model import AppMode
from models.workflow import Workflow, WorkflowKind
@@ -41,6 +43,7 @@ def test_run_uses_single_node_execution_branch(
app_config.app_id = "app"
app_config.tenant_id = "tenant"
app_config.workflow_id = "workflow"
app_config.app_mode = AppMode.WORKFLOW
app_generate_entity = MagicMock(spec=WorkflowAppGenerateEntity)
app_generate_entity.app_config = app_config
@@ -104,6 +107,7 @@ def test_run_uses_single_node_execution_branch(
single_iteration_run=single_iteration_run,
single_loop_run=single_loop_run,
user_id="user",
app_type=CreditUsageAppType.WORKFLOW,
trace_session_id="session-1",
)
init_graph.assert_not_called()
@@ -7,14 +7,19 @@ import pytest
from core.app.app_config.entities import WorkflowUIBasedAppConfig
from core.app.entities.app_invoke_entities import (
AdvancedChatAppGenerateEntity,
CreditUsageCreatedBy,
DifyRunContext,
InvokeFrom,
WorkflowAppGenerateEntity,
get_credit_usage_app_type,
get_credit_usage_created_by,
)
from core.app.layers.pause_state_persist_layer import (
WorkflowResumptionContext,
_AdvancedChatAppGenerateEntityWrapper,
_WorkflowGenerateEntityWrapper,
)
from core.credit_usage import CreditUsageAppType
from core.ops.ops_trace_manager import TraceQueueManager
from models.model import AppMode
@@ -99,6 +104,60 @@ def test_advanced_chat_generate_entity_roundtrip_excludes_trace_manager():
assert restored.trace_manager is None
@pytest.mark.parametrize(
("app_mode", "created_by"),
[
(AppMode.CHAT, CreditUsageCreatedBy.APP),
(AppMode.ADVANCED_CHAT, CreditUsageCreatedBy.APP),
(AppMode.WORKFLOW, CreditUsageCreatedBy.APP),
(AppMode.AGENT_CHAT, CreditUsageCreatedBy.APP),
(AppMode.AGENT, CreditUsageCreatedBy.APP),
(AppMode.COMPLETION, CreditUsageCreatedBy.APP),
(AppMode.RAG_PIPELINE, CreditUsageCreatedBy.APP),
],
)
def test_get_credit_usage_created_by_maps_app_mode(
app_mode: AppMode,
created_by: CreditUsageCreatedBy,
) -> None:
assert get_credit_usage_created_by(app_mode) == created_by
def test_get_credit_usage_created_by_defaults_to_unknown() -> None:
assert get_credit_usage_created_by(None) == CreditUsageCreatedBy.UNKNOWN
@pytest.mark.parametrize(
("app_mode", "app_type"),
[
(AppMode.CHAT, CreditUsageAppType.CHATBOT),
(AppMode.ADVANCED_CHAT, CreditUsageAppType.CHATFLOW),
(AppMode.WORKFLOW, CreditUsageAppType.WORKFLOW),
(AppMode.AGENT_CHAT, CreditUsageAppType.AGENT),
(AppMode.AGENT, CreditUsageAppType.AGENT_V2),
(AppMode.COMPLETION, CreditUsageAppType.COMPLETION),
(AppMode.RAG_PIPELINE, CreditUsageAppType.RAG_PIPELINE),
],
)
def test_get_credit_usage_app_type_maps_app_mode(app_mode: AppMode, app_type: CreditUsageAppType) -> None:
assert get_credit_usage_app_type(app_mode) == app_type
def test_dify_run_context_normalizes_unknown_created_by_values() -> None:
context = DifyRunContext(
tenant_id="tenant-id",
app_id="app-id",
user_id="user-id",
user_from="account",
invoke_from=InvokeFrom.DEBUGGER,
app_type="legacy-app-type",
created_by="legacy-created-by",
)
assert context.app_type == CreditUsageAppType.UNKNOWN
assert context.created_by == CreditUsageCreatedBy.UNKNOWN
@dataclass(frozen=True)
class ResumptionContextCase:
name: str
@@ -1043,7 +1043,7 @@ class TestEasyUiBasedGenerateTaskPipeline:
_set_method(pipeline, "_process_stream_response", lambda publisher, trace_manager: iter([payload]))
monkeypatch.setattr(
"core.app.task_pipeline.easy_ui_based_generate_task_pipeline.AppGeneratorTTSPublisher",
lambda tenant_id, voice, language: _Publisher(),
lambda tenant_id, voice, language, app_type: _Publisher(),
)
responses = list(pipeline._wrapper_process_stream_response())
@@ -1085,7 +1085,7 @@ class TestEasyUiBasedGenerateTaskPipeline:
_set_method(pipeline, "_process_stream_response", lambda publisher, trace_manager: iter([]))
monkeypatch.setattr(
"core.app.task_pipeline.easy_ui_based_generate_task_pipeline.AppGeneratorTTSPublisher",
lambda tenant_id, voice, language: _Publisher(),
lambda tenant_id, voice, language, app_type: _Publisher(),
)
responses = list(pipeline._wrapper_process_stream_response())
@@ -1124,7 +1124,7 @@ class TestEasyUiBasedGenerateTaskPipeline:
_set_method(pipeline, "_process_stream_response", lambda publisher, trace_manager: iter([]))
monkeypatch.setattr(
"core.app.task_pipeline.easy_ui_based_generate_task_pipeline.AppGeneratorTTSPublisher",
lambda tenant_id, voice, language: _Publisher(),
lambda tenant_id, voice, language, app_type: _Publisher(),
)
responses = list(pipeline._wrapper_process_stream_response())
@@ -1153,7 +1153,7 @@ class TestEasyUiBasedGenerateTaskPipeline:
_set_method(pipeline, "_process_stream_response", lambda publisher, trace_manager: iter([error]))
monkeypatch.setattr(
"core.app.task_pipeline.easy_ui_based_generate_task_pipeline.AppGeneratorTTSPublisher",
lambda tenant_id, voice, language: publisher,
lambda tenant_id, voice, language, app_type: publisher,
)
responses = list(pipeline._wrapper_process_stream_response())
@@ -10,6 +10,7 @@ from sqlalchemy.engine import Engine
from sqlalchemy.orm import sessionmaker
from configs import dify_config
from core.app.entities.app_invoke_entities import CreditUsageCreatedBy
from core.app.llm.quota import (
LLMQuotaReservationState,
deduct_llm_quota,
@@ -19,6 +20,7 @@ from core.app.llm.quota import (
reserve_llm_quota_for_model,
reserve_model_quota_for_model,
)
from core.credit_usage import CreditUsageAppType
from core.entities.model_entities import ModelStatus
from core.entities.provider_entities import ProviderQuotaType, QuotaUnit
from core.errors.error import QuotaExceededError
@@ -133,6 +135,8 @@ def test_reserve_llm_quota_uses_exact_credit_pool_reservation() -> None:
provider="openai",
model="gpt-4o",
request_id="11111111-1111-5111-8111-111111111111",
app_type=CreditUsageAppType.CHATBOT,
created_by=CreditUsageCreatedBy.APP.value,
)
reservation.commit(LLMUsage.empty_usage())
reservation.release()
@@ -145,7 +149,13 @@ def test_reserve_llm_quota_uses_exact_credit_pool_reservation() -> None:
pool_type="trial",
request_id="11111111-1111-5111-8111-111111111111",
session_factory=ANY,
meta={"source": "llm.invoke", "provider": "openai", "model": "gpt-4o"},
meta={
"source": "llm.invoke",
"provider": "openai",
"model": "gpt-4o",
"app_type": CreditUsageAppType.CHATBOT,
"created_by": CreditUsageCreatedBy.APP.value,
},
)
credit_reservation.commit.assert_called_once_with()
credit_reservation.release.assert_not_called()
@@ -227,6 +237,8 @@ def test_reserve_non_llm_quota_uses_model_type_and_credit_pool_reservation() ->
"provider": "openai",
"model_type": "text-embedding",
"model": "text-embedding-3-small",
"app_type": CreditUsageAppType.UNKNOWN,
"created_by": "unknown",
},
)
credit_reservation.commit.assert_called_once_with()
@@ -354,6 +366,13 @@ def test_deduct_llm_quota_for_model_uses_identity_based_trial_billing() -> None:
mock_deduct_credits.assert_called_once_with(
tenant_id="tenant-id",
credits_required=42,
metadata={
"provider": "openai",
"model": "gpt-4o",
"model_type": "llm",
"app_type": CreditUsageAppType.UNKNOWN,
"created_by": "unknown",
},
session=ANY,
)
@@ -474,6 +493,13 @@ def test_deduct_llm_quota_for_model_uses_credit_configuration() -> None:
mock_deduct_credits.assert_called_once_with(
tenant_id="tenant-id",
credits_required=9,
metadata={
"provider": "openai",
"model": "gpt-4o",
"model_type": "llm",
"app_type": CreditUsageAppType.UNKNOWN,
"created_by": "unknown",
},
session=ANY,
)
@@ -510,6 +536,13 @@ def test_deduct_llm_quota_for_model_uses_single_charge_for_times_quota() -> None
mock_deduct_credits.assert_called_once_with(
tenant_id="tenant-id",
credits_required=1,
metadata={
"provider": "openai",
"model": "gpt-4o",
"model_type": "llm",
"app_type": CreditUsageAppType.UNKNOWN,
"created_by": "unknown",
},
session=ANY,
)
@@ -548,6 +581,13 @@ def test_deduct_llm_quota_for_model_uses_paid_billing_pool() -> None:
tenant_id="tenant-id",
credits_required=5,
pool_type="paid",
metadata={
"provider": "openai",
"model": "gpt-4o",
"model_type": "llm",
"app_type": CreditUsageAppType.UNKNOWN,
"created_by": "unknown",
},
session=ANY,
)
@@ -7,12 +7,14 @@ from core.moderation.base import ModerationAction, ModerationError, ModerationIn
from core.moderation.input_moderation import InputModeration
from core.ops.entities.trace_entity import TraceTaskName
from core.ops.ops_trace_manager import TraceQueueManager
from models.model import AppMode
class TestInputModeration:
@pytest.fixture
def app_config(self):
config = MagicMock(spec=AppConfig)
config.app_mode = AppMode.CHAT
config.sensitive_word_avoidance = None
return config
@@ -755,7 +755,7 @@ class TestIndexingRunnerLoad:
# Verify executor was used for parallel processing
assert mock_executor_instance.submit.called
for submit_call in mock_executor_instance.submit.call_args_list:
assert submit_call.args[0] == runner._process_chunk
assert submit_call.args[0].__name__ == runner._process_chunk.__name__
assert len(submit_call.args) == 6
mock_future.result.assert_called()
assert mock_update_status.call_args.kwargs["extra_update_params"][DatasetDocument.tokens] == 300
@@ -0,0 +1,33 @@
from concurrent.futures import ThreadPoolExecutor
from core.credit_usage import CreditUsageAppType, CreditUsageCreatedBy
from core.model_context import get_credit_usage_metadata, use_credit_usage_metadata, with_credit_usage_created_by
from extensions.otel import propagate_context
def test_credit_usage_metadata_is_scoped_and_propagated_to_worker_threads() -> None:
with use_credit_usage_metadata({"app_type": CreditUsageAppType.WORKFLOW}):
assert get_credit_usage_metadata() == {"app_type": CreditUsageAppType.WORKFLOW}
with ThreadPoolExecutor(max_workers=1) as executor:
read_context = propagate_context(get_credit_usage_metadata)
assert executor.submit(read_context).result() == {"app_type": CreditUsageAppType.WORKFLOW}
assert get_credit_usage_metadata() is None
def test_nested_credit_usage_metadata_keeps_app_type_and_direct_feature() -> None:
with use_credit_usage_metadata({"app_type": CreditUsageAppType.WORKFLOW}):
with use_credit_usage_metadata({"created_by": CreditUsageCreatedBy.KNOWLEDGE_RETRIEVAL}):
assert get_credit_usage_metadata() == {
"app_type": CreditUsageAppType.WORKFLOW,
"created_by": CreditUsageCreatedBy.KNOWLEDGE_RETRIEVAL,
}
def test_credit_usage_created_by_decorator_sets_feature() -> None:
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_RETRIEVAL)
def read_context() -> object:
return get_credit_usage_metadata()
assert read_context() == {"created_by": CreditUsageCreatedBy.KNOWLEDGE_RETRIEVAL}
@@ -7,6 +7,7 @@ import pytest
import redis
from pytest_mock import MockerFixture
from core.app.entities.app_invoke_entities import CreditUsageCreatedBy
from core.entities.provider_entities import (
ModelLoadBalancingConfiguration,
ProviderQuotaType,
@@ -184,11 +185,14 @@ def test_quota_managed_non_streaming_invocation_finalizes_reservation() -> None:
response = model_instance.invoke_llm(
prompt_messages=[],
stream=False,
request_metadata={"invocation_id": invocation_id},
request_metadata={"invocation_id": invocation_id, "created_by": CreditUsageCreatedBy.APP.value},
)
assert response is result
reserve_quota.assert_called_once_with(request_id=invocation_id)
reserve_quota.assert_called_once_with(
request_id=invocation_id,
created_by=CreditUsageCreatedBy.APP.value,
)
invoke.assert_called_once()
reservation.commit.assert_called_once_with(usage)
reservation.release.assert_called_once_with()
@@ -481,6 +481,8 @@ class TestDifyNodeFactoryCreateNode:
app_id="app-id",
user_id="user-id",
invoke_from=InvokeFrom.DEBUGGER,
app_type=None,
created_by=None,
)
factory._code_executor = sentinel.code_executor
factory._code_limits = sentinel.code_limits
@@ -277,7 +277,7 @@ def test_dify_prepared_llm_requires_model_schema() -> None:
def test_dify_prepared_llm_delegates_structured_output_helper(monkeypatch: pytest.MonkeyPatch) -> None:
model_instance = _ModelInstanceStub(model_schema=_build_model_schema())
prepared = DifyPreparedLLM(model_instance)
prepared = DifyPreparedLLM(model_instance, request_metadata={"created_by": "workflow"})
invoke_structured = MagicMock(return_value=sentinel.structured)
monkeypatch.setattr(node_runtime, "invoke_llm_with_structured_output", invoke_structured)
@@ -299,6 +299,7 @@ def test_dify_prepared_llm_delegates_structured_output_helper(monkeypatch: pytes
model_parameters={"temperature": 0.2},
stop=["done"],
stream=True,
request_metadata={"created_by": "workflow"},
)
@@ -8,6 +8,7 @@ import pytest
from core.app.apps.exc import GenerateTaskStoppedError
from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom
from core.credit_usage import CreditUsageAppType
from core.workflow import workflow_entry
from core.workflow.system_variables import default_system_variables
from graphon.entities.base_node_data import BaseNodeData
@@ -610,6 +611,7 @@ class TestWorkflowEntryHelpers:
user_id="user-id",
user_from=UserFrom.ACCOUNT,
invoke_from=InvokeFrom.DEBUGGER,
app_type=CreditUsageAppType.WORKFLOW,
)
graph_init_context_cls.assert_called_once_with(
workflow_id="",
@@ -7,10 +7,11 @@ import pytest
from sqlalchemy import select
from sqlalchemy.orm import Session, sessionmaker
from core.app.entities.app_invoke_entities import AgentAppGenerateEntity, ChatAppGenerateEntity
from core.app.entities.app_invoke_entities import AgentAppGenerateEntity, ChatAppGenerateEntity, CreditUsageCreatedBy
from core.credit_usage import CreditUsageAppType
from core.entities.provider_entities import ProviderQuotaType, QuotaUnit
from events.event_handlers import update_provider_when_message_created
from models import Message, TenantCreditPool
from models import AppMode, Message, TenantCreditPool
from models.enums import ProviderQuotaType as ModelProviderQuotaType
from models.provider import ProviderType
@@ -48,7 +49,7 @@ def test_message_created_trial_credit_accounting_does_not_raise_when_balance_is_
],
)
application_generate_entity = ChatAppGenerateEntity.model_construct(
app_config=SimpleNamespace(tenant_id=tenant_id),
app_config=SimpleNamespace(tenant_id=tenant_id, app_mode=AppMode.CHAT),
model_conf=SimpleNamespace(
provider="openai",
model="gpt-4o",
@@ -61,6 +62,7 @@ def test_message_created_trial_credit_accounting_does_not_raise_when_balance_is_
),
)
message = Message(message_tokens=2, answer_tokens=1)
message.id = "message-1"
with (
patch.object(update_provider_when_message_created, "_execute_provider_updates"),
@@ -89,7 +91,7 @@ def test_message_created_paid_credit_accounting_uses_paid_pool() -> None:
],
)
application_generate_entity = ChatAppGenerateEntity.model_construct(
app_config=SimpleNamespace(tenant_id=tenant_id),
app_config=SimpleNamespace(tenant_id=tenant_id, app_mode=AppMode.CHAT),
model_conf=SimpleNamespace(
provider="openai",
model="gpt-4o",
@@ -102,6 +104,7 @@ def test_message_created_paid_credit_accounting_uses_paid_pool() -> None:
),
)
message = Message(message_tokens=2, answer_tokens=1)
message.id = "message-1"
with (
patch.object(update_provider_when_message_created, "_deduct_credit_pool_quota_capped") as mock_deduct,
@@ -116,6 +119,14 @@ def test_message_created_paid_credit_accounting_uses_paid_pool() -> None:
tenant_id=tenant_id,
credits_required=3,
pool_type="paid",
request_id="message-1",
metadata={
"provider": "openai",
"model": "gpt-4o",
"model_type": "llm",
"app_type": CreditUsageAppType.CHATBOT,
"created_by": CreditUsageCreatedBy.APP.value,
},
)
@@ -132,7 +143,7 @@ def test_agent_app_gateway_accounting_skips_legacy_message_charge() -> None:
],
)
application_generate_entity = AgentAppGenerateEntity.model_construct(
app_config=SimpleNamespace(tenant_id=tenant_id),
app_config=SimpleNamespace(tenant_id=tenant_id, app_mode=AppMode.AGENT),
model_conf=SimpleNamespace(
provider="openai",
model="gpt-4o",
@@ -11,6 +11,7 @@ from pydantic import JsonValue
from sqlalchemy.orm import Session, sessionmaker
from configs import dify_config
from core.credit_usage import CreditUsageAppType, CreditUsageCreatedBy
from core.entities.model_entities import ModelStatus
from core.entities.provider_entities import ProviderQuotaType, QuotaUnit
from core.model_manager import ModelInstance, QuotaManagedModelInstance
@@ -63,6 +64,7 @@ def _model_instance(
instance.credentials = {"api_key": "hosted"}
instance.model_type_instance = MagicMock()
instance.load_balancing_manager = None
instance._request_metadata = None
return instance, provider_model
@@ -71,13 +73,14 @@ def _persist_app(
*,
request: AgentLLMInvokeRequest,
tenant_id: str | None = None,
mode: AppMode = AppMode.CHAT,
) -> App:
app = App(
id=request.caller.app_id,
tenant_id=tenant_id or request.caller.tenant_id,
name="Agent LLM gateway test app",
description="",
mode=AppMode.CHAT,
mode=mode,
enable_site=False,
enable_api=False,
max_active_requests=None,
@@ -175,7 +178,7 @@ def test_gateway_uses_quota_managed_instance_as_single_credit_owner(
sqlite_session: Session,
) -> None:
request = _request()
_persist_app(sqlite_session, request=request)
_persist_app(sqlite_session, request=request, mode=AppMode.WORKFLOW)
service = AgentLLMInnerService(session_factory=sqlite_session_factory)
model_instance, _ = _model_instance()
reservation = MagicMock(commit_before_delivery=True)
@@ -190,7 +193,12 @@ 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(request_id=request.caller.invocation_id)
model_instance.reserve_quota.assert_called_once_with(
request_id=request.caller.invocation_id,
app_type=CreditUsageAppType.WORKFLOW,
created_by=CreditUsageCreatedBy.AGENT_NODE,
)
assert provider_invoke.call_args.kwargs["request_metadata"]["agent_config_version_kind"] == "draft"
reservation.commit.assert_called_once_with(provider_chunk.delta.usage)
reservation.release.assert_called_once_with()
provider_invoke.assert_called_once()
@@ -227,6 +235,39 @@ def test_gateway_forwards_prompt_messages_without_revalidation() -> None:
assert model_instance.invoke_llm.call_args.kwargs["prompt_messages"] is prompt_messages
def test_standalone_agent_app_gateway_is_attributed_to_app() -> None:
request = _request()
model_instance = MagicMock(spec=ModelInstance)
model_instance.invoke_llm.return_value = iter([])
prepared = PreparedAgentLLMInvocation(
request=request,
model_instance=model_instance,
app_type=CreditUsageAppType.AGENT_V2,
)
assert list(AgentLLMInnerService().invoke(prepared)) == []
request_metadata = model_instance.invoke_llm.call_args.kwargs["request_metadata"]
assert request_metadata["created_by"] is CreditUsageCreatedBy.APP
assert request_metadata["agent_config_version_kind"] == "draft"
def test_agent_build_draft_gateway_is_attributed_to_build_draft() -> None:
request = _request()
request.caller.agent_config_version_kind = "build_draft"
model_instance = MagicMock(spec=ModelInstance)
model_instance.invoke_llm.return_value = iter([])
prepared = PreparedAgentLLMInvocation(
request=request,
model_instance=model_instance,
app_type=CreditUsageAppType.AGENT_V2,
)
assert list(AgentLLMInnerService().invoke(prepared)) == []
request_metadata = model_instance.invoke_llm.call_args.kwargs["request_metadata"]
assert request_metadata["created_by"] is CreditUsageCreatedBy.BUILD_DRAFT
assert request_metadata["agent_config_version_kind"] == "build_draft"
def test_retried_gateway_delivery_uses_one_effective_billing_charge(
sqlite_session_factory: sessionmaker[Session],
sqlite_session: Session,
@@ -64,6 +64,7 @@ from flask import Flask
from sqlalchemy.orm import Session
from werkzeug.datastructures import FileStorage
from core.credit_usage import CreditUsageAppType, CreditUsageCreatedBy
from core.plugin.entities.plugin_daemon import TTSAudioChunk
from graphon.model_runtime.errors.invoke import InvokeBadRequestError
from models.agent_config_entities import AgentSoulConfig
@@ -282,7 +283,14 @@ class TestAudioServiceASR:
# Assert
assert result == {"text": "Transcribed text"}
mock_model_instance.invoke_speech2text.assert_called_once()
mock_model_manager_class.assert_called_once_with(tenant_id=app.tenant_id, user_id="user-123")
mock_model_manager_class.assert_called_once_with(
tenant_id=app.tenant_id,
user_id="user-123",
request_metadata={
"app_type": CreditUsageAppType.CHATBOT,
"created_by": CreditUsageCreatedBy.AUDIO,
},
)
@patch("services.audio_service.ModelManager.for_tenant", autospec=True)
def test_transcript_asr_accepts_x_m4a_mimetype(
@@ -393,7 +401,14 @@ class TestAudioServiceASR:
)
assert result == {"text": "Agent transcript"}
mock_model_manager_class.assert_called_once_with(tenant_id=app.tenant_id, user_id="account-1")
mock_model_manager_class.assert_called_once_with(
tenant_id=app.tenant_id,
user_id="account-1",
request_metadata={
"app_type": CreditUsageAppType.AGENT_V2,
"created_by": CreditUsageCreatedBy.AUDIO,
},
)
@pytest.mark.parametrize(
"agent_soul",
@@ -586,7 +601,14 @@ class TestAudioServiceTTS:
# Assert
assert result.content_type == "audio/mpeg"
assert result.get_data() == b"audio data"
mock_model_manager_class.assert_called_once_with(tenant_id=app.tenant_id, user_id="user-123")
mock_model_manager_class.assert_called_once_with(
tenant_id=app.tenant_id,
user_id="user-123",
request_metadata={
"app_type": CreditUsageAppType.CHATBOT,
"created_by": CreditUsageCreatedBy.AUDIO,
},
)
mock_model_instance.invoke_tts.assert_called_once_with(
content_text="Hello world",
voice="en-US-Neural",
@@ -9,6 +9,8 @@ import pytest
from sqlalchemy import select
from sqlalchemy.orm import Session
from core.app.entities.app_invoke_entities import CreditUsageCreatedBy
from core.credit_usage import CreditUsageAppType
from core.errors.error import QuotaExceededError
from enums import DeploymentEdition
from models import TenantCreditPool
@@ -301,7 +303,11 @@ def test_reserve_credits_commits_billing_reservation_once() -> None:
bucket="trial",
request_id="request-1",
amount=3,
meta={"source": "test"},
meta={
"source": "test",
"created_by": CreditUsageCreatedBy.UNKNOWN.value,
"app_type": CreditUsageAppType.UNKNOWN.value,
},
)
quota_commit.assert_called_once_with(
tenant_id="tenant-1",
@@ -309,7 +315,12 @@ def test_reserve_credits_commits_billing_reservation_once() -> None:
bucket="trial",
reservation_id="reservation-1",
actual_amount=3,
meta={"source": "test", "request_id": "request-1"},
meta={
"source": "test",
"created_by": CreditUsageCreatedBy.UNKNOWN.value,
"app_type": CreditUsageAppType.UNKNOWN.value,
"request_id": "request-1",
},
)
quota_release.assert_not_called()
@@ -381,7 +392,11 @@ def test_check_and_deduct_credits_uses_billing_reserve_and_commit_when_enabled()
bucket="trial",
request_id=ANY,
amount=3,
meta={"source": "credit_pool.check_and_deduct"},
meta={
"source": "credit_pool.check_and_deduct",
"created_by": CreditUsageCreatedBy.UNKNOWN.value,
"app_type": CreditUsageAppType.UNKNOWN.value,
},
)
quota_commit.assert_called_once_with(
tenant_id=tenant_id,
@@ -389,7 +404,11 @@ def test_check_and_deduct_credits_uses_billing_reserve_and_commit_when_enabled()
bucket="trial",
reservation_id="reservation-1",
actual_amount=3,
meta={"source": "credit_pool.check_and_deduct"},
meta={
"source": "credit_pool.check_and_deduct",
"created_by": CreditUsageCreatedBy.UNKNOWN.value,
"app_type": CreditUsageAppType.UNKNOWN.value,
},
)
quota_release.assert_not_called()
@@ -413,6 +432,8 @@ def test_check_and_deduct_credits_forwards_deterministic_billing_identity() -> N
assert result == 3
expected_metadata = {
"source": "credit_pool.check_and_deduct",
"created_by": CreditUsageCreatedBy.UNKNOWN.value,
"app_type": CreditUsageAppType.UNKNOWN.value,
"agent_run_id": "run-1",
}
quota_reserve.assert_called_once_with(
@@ -509,6 +530,13 @@ def test_deduct_credits_capped_uses_billing_consume_capped_when_enabled() -> Non
tenant_id=tenant_id,
credits_required=5,
pool_type=ProviderQuotaType.PAID,
request_id="message-1",
metadata={
"provider": "openai",
"model": "gpt-4o",
"app_type": CreditUsageAppType.CHATBOT,
"created_by": CreditUsageCreatedBy.APP,
},
)
assert result == 2
@@ -516,9 +544,15 @@ def test_deduct_credits_capped_uses_billing_consume_capped_when_enabled() -> Non
tenant_id=tenant_id,
feature_key=FEATURE_KEY_CREDIT_POOL,
bucket="paid",
request_id=ANY,
request_id="message-1",
amount=5,
meta={"source": "credit_pool.deduct_capped"},
meta={
"source": "credit_pool.deduct_capped",
"provider": "openai",
"model": "gpt-4o",
"app_type": CreditUsageAppType.CHATBOT.value,
"created_by": CreditUsageCreatedBy.APP.value,
},
)
@@ -10,6 +10,7 @@ stay fast and focus on the wiring itself.
from unittest.mock import MagicMock, patch
from core.app.app_config.entities import ModelConfig
from core.credit_usage import CreditUsageAppType, CreditUsageCreatedBy
from graphon.model_runtime.entities.llm_entities import LLMMode
from services.workflow_generator_service import WorkflowGeneratorService
@@ -63,7 +64,13 @@ class TestWorkflowGeneratorService:
)
# Assert
mock_model_manager.for_tenant.assert_called_once_with(tenant_id="t-1")
mock_model_manager.for_tenant.assert_called_once_with(
tenant_id="t-1",
request_metadata={
"app_type": CreditUsageAppType.WORKFLOW,
"created_by": CreditUsageCreatedBy.WORKFLOW_GENERATION,
},
)
mock_workflow_generator.generate_workflow_graph.assert_called_once()
call_kwargs = mock_workflow_generator.generate_workflow_graph.call_args.kwargs
assert call_kwargs["model_instance"] is instance