From 960018253b2a108f46c630485c83faf1203ccce0 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=9E=97=E7=8E=AE=20=28Jade=20Lin=29?= Date: Tue, 25 Aug 2026 10:21:32 +0000 Subject: [PATCH] feat(api): track credit usage contexts (#41181) Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> --- api/core/agent/cot_agent_runner.py | 33 ++- api/core/agent/fc_agent_runner.py | 39 ++-- api/core/app/apps/advanced_chat/app_runner.py | 5 + .../advanced_chat/generate_task_pipeline.py | 6 +- api/core/app/apps/agent_app/app_generator.py | 2 + api/core/app/apps/chat/app_runner.py | 8 +- api/core/app/apps/completion/app_runner.py | 8 +- api/core/app/apps/pipeline/pipeline_runner.py | 2 + api/core/app/apps/workflow/app_runner.py | 11 +- .../apps/workflow/generate_task_pipeline.py | 7 +- api/core/app/apps/workflow_app_runner.py | 8 + api/core/app/entities/app_invoke_entities.py | 58 +++++ api/core/app/llm/quota.py | 75 ++++++- .../easy_ui_based_generate_task_pipeline.py | 6 +- .../base/tts/app_generator_tts_publisher.py | 16 +- api/core/credit_usage.py | 74 ++++++ api/core/indexing_runner.py | 11 +- api/core/llm_generator/llm_generator.py | 11 + .../output_parser/structured_output.py | 5 + api/core/model_context.py | 61 +++++ api/core/model_manager.py | 212 +++++++++++++++--- api/core/moderation/input_moderation.py | 13 +- .../openai_moderation/openai_moderation.py | 3 + api/core/plugin/backwards_invocation/model.py | 3 + .../processor/paragraph_index_processor.py | 3 + api/core/rag/rerank/rerank_model.py | 3 + api/core/rag/rerank/weight_rerank.py | 3 + api/core/rag/retrieval/dataset_retrieval.py | 26 ++- .../multi_dataset_function_call_router.py | 3 + .../router/multi_dataset_react_route.py | 3 + api/core/rag/summary_index/summary_index.py | 6 +- .../builtin_tool/providers/audio/tools/asr.py | 3 + .../builtin_tool/providers/audio/tools/tts.py | 3 + .../tools/utils/model_invocation_utils.py | 3 + api/core/workflow/node_factory.py | 12 +- api/core/workflow/node_runtime.py | 46 ++-- .../knowledge_index/knowledge_index_node.py | 21 +- .../knowledge_retrieval_node.py | 7 + .../nodes/knowledge_retrieval/retrieval.py | 3 + api/core/workflow/workflow_entry.py | 9 +- .../update_provider_when_message_created.py | 57 ++++- api/extensions/otel/context.py | 4 +- api/services/agent_llm_inner_service.py | 25 ++- api/services/audio_service.py | 25 ++- api/services/conversation_service.py | 8 +- api/services/credit_pool_service.py | 33 ++- api/services/message_service.py | 8 +- api/services/summary_index_service.py | 5 + api/services/vector_service.py | 8 + api/services/workflow_generator_service.py | 18 +- .../batch_create_segment_to_index_task.py | 3 + api/tasks/regenerate_summary_index_task.py | 3 + .../core/agent/test_cot_agent_runner.py | 8 +- .../core/agent/test_fc_agent_runner.py | 8 +- .../chat/test_app_generator_and_runner.py | 3 + .../app/apps/completion/test_app_runner.py | 8 +- .../test_workflow_app_runner_single_node.py | 4 + .../app/entities/test_app_invoke_entities.py | 59 +++++ ...sy_ui_based_generate_task_pipeline_core.py | 8 +- .../unit_tests/core/app/test_llm_quota.py | 42 +++- .../core/moderation/test_input_moderation.py | 2 + .../core/rag/indexing/test_indexing_runner.py | 2 +- .../unit_tests/core/test_model_context.py | 33 +++ .../unit_tests/core/test_model_manager.py | 8 +- .../core/workflow/test_node_factory.py | 2 + .../core/workflow/test_node_runtime.py | 3 +- .../workflow/test_workflow_entry_helpers.py | 2 + ...st_update_provider_when_message_created.py | 21 +- .../services/test_agent_llm_inner_service.py | 47 +++- .../unit_tests/services/test_audio_service.py | 28 ++- .../services/test_credit_pool_service.py | 46 +++- .../test_workflow_generator_service.py | 9 +- 72 files changed, 1206 insertions(+), 165 deletions(-) create mode 100644 api/core/credit_usage.py create mode 100644 api/core/model_context.py create mode 100644 api/tests/unit_tests/core/test_model_context.py diff --git a/api/core/agent/cot_agent_runner.py b/api/core/agent/cot_agent_runner.py index a1ccf75386b..ca73c3f6671 100644 --- a/api/core/agent/cot_agent_runner.py +++ b/api/core/agent/cot_agent_runner.py @@ -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() diff --git a/api/core/agent/fc_agent_runner.py b/api/core/agent/fc_agent_runner.py index 78980f0d943..18b0d2ce016 100644 --- a/api/core/agent/fc_agent_runner.py +++ b/api/core/agent/fc_agent_runner.py @@ -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 diff --git a/api/core/app/apps/advanced_chat/app_runner.py b/api/core/app/apps/advanced_chat/app_runner.py index 4b303663e08..5c6d28f524c 100644 --- a/api/core/app/apps/advanced_chat/app_runner.py +++ b/api/core/app/apps/advanced_chat/app_runner.py @@ -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"), ) ) diff --git a/api/core/app/apps/advanced_chat/generate_task_pipeline.py b/api/core/app/apps/advanced_chat/generate_task_pipeline.py index 43c100bab32..b2ad5558e22 100644 --- a/api/core/app/apps/advanced_chat/generate_task_pipeline.py +++ b/api/core/app/apps/advanced_chat/generate_task_pipeline.py @@ -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: diff --git a/api/core/app/apps/agent_app/app_generator.py b/api/core/app/apps/agent_app/app_generator.py index 7d503d07d7a..b2457a44a04 100644 --- a/api/core/app/apps/agent_app/app_generator.py +++ b/api/core/app/apps/agent_app/app_generator.py @@ -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( diff --git a/api/core/app/apps/chat/app_runner.py b/api/core/app/apps/chat/app_runner.py index 25ade27e76f..65d60748d1e 100644 --- a/api/core/app/apps/chat/app_runner.py +++ b/api/core/app/apps/chat/app_runner.py @@ -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 diff --git a/api/core/app/apps/completion/app_runner.py b/api/core/app/apps/completion/app_runner.py index 9545dda62f5..25828e5dd25 100644 --- a/api/core/app/apps/completion/app_runner.py +++ b/api/core/app/apps/completion/app_runner.py @@ -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 diff --git a/api/core/app/apps/pipeline/pipeline_runner.py b/api/core/app/apps/pipeline/pipeline_runner.py index 39bd4c6980c..651b670fcce 100644 --- a/api/core/app/apps/pipeline/pipeline_runner.py +++ b/api/core/app/apps/pipeline/pipeline_runner.py @@ -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, diff --git a/api/core/app/apps/workflow/app_runner.py b/api/core/app/apps/workflow/app_runner.py index b2a3b9dd77a..df545d0e345 100644 --- a/api/core/app/apps/workflow/app_runner.py +++ b/api/core/app/apps/workflow/app_runner.py @@ -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"), ) ) diff --git a/api/core/app/apps/workflow/generate_task_pipeline.py b/api/core/app/apps/workflow/generate_task_pipeline.py index 6c13bcfbd39..ac73d951e3a 100644 --- a/api/core/app/apps/workflow/generate_task_pipeline.py +++ b/api/core/app/apps/workflow/generate_task_pipeline.py @@ -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: diff --git a/api/core/app/apps/workflow_app_runner.py b/api/core/app/apps/workflow_app_runner.py index f8cf6e4fc9c..7821e6c9fbc 100644 --- a/api/core/app/apps/workflow_app_runner.py +++ b/api/core/app/apps/workflow_app_runner.py @@ -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( diff --git a/api/core/app/entities/app_invoke_entities.py b/api/core/app/entities/app_invoke_entities.py index c59208c4aaf..7f983135994 100644 --- a/api/core/app/entities/app_invoke_entities.py +++ b/api/core/app/entities/app_invoke_entities.py @@ -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 diff --git a/api/core/app/llm/quota.py b/api/core/app/llm/quota.py index 2c892b9fd98..1417c9de0cf 100644 --- a/api/core/app/llm/quota.py +++ b/api/core/app/llm/quota.py @@ -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, ) diff --git a/api/core/app/task_pipeline/easy_ui_based_generate_task_pipeline.py b/api/core/app/task_pipeline/easy_ui_based_generate_task_pipeline.py index 283e18ab376..1a9989a2039 100644 --- a/api/core/app/task_pipeline/easy_ui_based_generate_task_pipeline.py +++ b/api/core/app/task_pipeline/easy_ui_based_generate_task_pipeline.py @@ -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): diff --git a/api/core/base/tts/app_generator_tts_publisher.py b/api/core/base/tts/app_generator_tts_publisher.py index c0ad76e2951..b0e1384ca6b 100644 --- a/api/core/base/tts/app_generator_tts_publisher.py +++ b/api/core/base/tts/app_generator_tts_publisher.py @@ -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 ) diff --git a/api/core/credit_usage.py b/api/core/credit_usage.py new file mode 100644 index 00000000000..1ea871c6464 --- /dev/null +++ b/api/core/credit_usage.py @@ -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 diff --git a/api/core/indexing_runner.py b/api/core/indexing_runner.py index 7526bc533cc..02b9255cfc6 100644 --- a/api/core/indexing_runner.py +++ b/api/core/indexing_runner.py @@ -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, diff --git a/api/core/llm_generator/llm_generator.py b/api/core/llm_generator/llm_generator.py index e8b1b023687..c6842b2b86c 100644 --- a/api/core/llm_generator/llm_generator.py +++ b/api/core/llm_generator/llm_generator.py @@ -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, diff --git a/api/core/llm_generator/output_parser/structured_output.py b/api/core/llm_generator/output_parser/structured_output.py index f2e98244197..a48b8da957f 100644 --- a/api/core/llm_generator/output_parser/structured_output.py +++ b/api/core/llm_generator/output_parser/structured_output.py @@ -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): diff --git a/api/core/model_context.py b/api/core/model_context.py new file mode 100644 index 00000000000..48ecc384cc9 --- /dev/null +++ b/api/core/model_context.py @@ -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 diff --git a/api/core/model_manager.py b/api/core/model_manager.py index 81ef887d6dc..4d197146a1d 100644 --- a/api/core/model_manager.py +++ b/api/core/model_manager.py @@ -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 diff --git a/api/core/moderation/input_moderation.py b/api/core/moderation/input_moderation.py index 21dc58f16f4..9edcef3be38 100644 --- a/api/core/moderation/input_moderation.py +++ b/api/core/moderation/input_moderation.py @@ -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( diff --git a/api/core/moderation/openai_moderation/openai_moderation.py b/api/core/moderation/openai_moderation/openai_moderation.py index 4b7a08eb277..5b0da9eedeb 100644 --- a/api/core/moderation/openai_moderation/openai_moderation.py +++ b/api/core/moderation/openai_moderation/openai_moderation.py @@ -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) diff --git a/api/core/plugin/backwards_invocation/model.py b/api/core/plugin/backwards_invocation/model.py index 193490673e0..be8f5f8d7ac 100644 --- a/api/core/plugin/backwards_invocation/model.py +++ b/api/core/plugin/backwards_invocation/model.py @@ -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, diff --git a/api/core/rag/index_processor/processor/paragraph_index_processor.py b/api/core/rag/index_processor/processor/paragraph_index_processor.py index d6323039c38..d0b322db70c 100644 --- a/api/core/rag/index_processor/processor/paragraph_index_processor.py +++ b/api/core/rag/index_processor/processor/paragraph_index_processor.py @@ -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, diff --git a/api/core/rag/rerank/rerank_model.py b/api/core/rag/rerank/rerank_model.py index 3dc0517860f..1857117fbd6 100644 --- a/api/core/rag/rerank/rerank_model.py +++ b/api/core/rag/rerank/rerank_model.py @@ -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, diff --git a/api/core/rag/rerank/weight_rerank.py b/api/core/rag/rerank/weight_rerank.py index 3743d98fb98..49cb7926983 100644 --- a/api/core/rag/rerank/weight_rerank.py +++ b/api/core/rag/rerank/weight_rerank.py @@ -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]: diff --git a/api/core/rag/retrieval/dataset_retrieval.py b/api/core/rag/retrieval/dataset_retrieval.py index 197ad138360..a9a3cbb16b9 100644 --- a/api/core/rag/retrieval/dataset_retrieval.py +++ b/api/core/rag/retrieval/dataset_retrieval.py @@ -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, diff --git a/api/core/rag/retrieval/router/multi_dataset_function_call_router.py b/api/core/rag/retrieval/router/multi_dataset_function_call_router.py index ef56567a5a9..3dc6a6e1b2c 100644 --- a/api/core/rag/retrieval/router/multi_dataset_function_call_router.py +++ b/api/core/rag/retrieval/router/multi_dataset_function_call_router.py @@ -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, diff --git a/api/core/rag/retrieval/router/multi_dataset_react_route.py b/api/core/rag/retrieval/router/multi_dataset_react_route.py index 95ffd4c84ca..2c7237825d7 100644 --- a/api/core/rag/retrieval/router/multi_dataset_react_route.py +++ b/api/core/rag/retrieval/router/multi_dataset_react_route.py @@ -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], diff --git a/api/core/rag/summary_index/summary_index.py b/api/core/rag/summary_index/summary_index.py index d9ce3879890..89788c6ab53 100644 --- a/api/core/rag/summary_index/summary_index.py +++ b/api/core/rag/summary_index/summary_index.py @@ -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) diff --git a/api/core/tools/builtin_tool/providers/audio/tools/asr.py b/api/core/tools/builtin_tool/providers/audio/tools/asr.py index 8d5959d2d1b..c315764c3b3 100644 --- a/api/core/tools/builtin_tool/providers/audio/tools/asr.py +++ b/api/core/tools/builtin_tool/providers/audio/tools/asr.py @@ -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, diff --git a/api/core/tools/builtin_tool/providers/audio/tools/tts.py b/api/core/tools/builtin_tool/providers/audio/tools/tts.py index 7f15f42d679..a8083e74e25 100644 --- a/api/core/tools/builtin_tool/providers/audio/tools/tts.py +++ b/api/core/tools/builtin_tool/providers/audio/tools/tts.py @@ -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, diff --git a/api/core/tools/utils/model_invocation_utils.py b/api/core/tools/utils/model_invocation_utils.py index a3623d4ecde..901ce24f966 100644 --- a/api/core/tools/utils/model_invocation_utils.py +++ b/api/core/tools/utils/model_invocation_utils.py @@ -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, diff --git a/api/core/workflow/node_factory.py b/api/core/workflow/node_factory.py index bedbb5765c4..1d868c58f9c 100644 --- a/api/core/workflow/node_factory.py +++ b/api/core/workflow/node_factory.py @@ -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 diff --git a/api/core/workflow/node_runtime.py b/api/core/workflow/node_runtime.py index f113a88ccb9..57315d7d6c3 100644 --- a/api/core/workflow/node_runtime.py +++ b/api/core/workflow/node_runtime.py @@ -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 diff --git a/api/core/workflow/nodes/knowledge_index/knowledge_index_node.py b/api/core/workflow/nodes/knowledge_index/knowledge_index_node.py index e0721a15ee5..944a6bb9109 100644 --- a/api/core/workflow/nodes/knowledge_index/knowledge_index_node.py +++ b/api/core/workflow/nodes/knowledge_index/knowledge_index_node.py @@ -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 diff --git a/api/core/workflow/nodes/knowledge_retrieval/knowledge_retrieval_node.py b/api/core/workflow/nodes/knowledge_retrieval/knowledge_retrieval_node.py index b4e975dcfa6..a3547fbcf76 100644 --- a/api/core/workflow/nodes/knowledge_retrieval/knowledge_retrieval_node.py +++ b/api/core/workflow/nodes/knowledge_retrieval/knowledge_retrieval_node.py @@ -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") diff --git a/api/core/workflow/nodes/knowledge_retrieval/retrieval.py b/api/core/workflow/nodes/knowledge_retrieval/retrieval.py index ea45dcf5c20..c588a37d698 100644 --- a/api/core/workflow/nodes/knowledge_retrieval/retrieval.py +++ b/api/core/workflow/nodes/knowledge_retrieval/retrieval.py @@ -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]: ... diff --git a/api/core/workflow/workflow_entry.py b/api/core/workflow/workflow_entry.py index 372bbd4e7f8..4ce22fa23d8 100644 --- a/api/core/workflow/workflow_entry.py +++ b/api/core/workflow/workflow_entry.py @@ -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="", diff --git a/api/events/event_handlers/update_provider_when_message_created.py b/api/events/event_handlers/update_provider_when_message_created.py index a2ed5568c99..af60032d099 100644 --- a/api/events/event_handlers/update_provider_when_message_created.py +++ b/api/events/event_handlers/update_provider_when_message_created.py @@ -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, " diff --git a/api/extensions/otel/context.py b/api/extensions/otel/context.py index b7378a3e667..3426e2dcf08 100644 --- a/api/extensions/otel/context.py +++ b/api/extensions/otel/context.py @@ -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) diff --git a/api/services/agent_llm_inner_service.py b/api/services/agent_llm_inner_service.py index 8e4fd80f24e..4e39f358ba8 100644 --- a/api/services/agent_llm_inner_service.py +++ b/api/services/agent_llm_inner_service.py @@ -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"] diff --git a/api/services/audio_service.py b/api/services/audio_service.py index 7899fa15d33..9b97ee2fe3f 100644 --- a/api/services/audio_service.py +++ b/api/services/audio_service.py @@ -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() diff --git a/api/services/conversation_service.py b/api/services/conversation_service.py index e4622056d28..68e8d3c3e2c 100644 --- a/api/services/conversation_service.py +++ b/api/services/conversation_service.py @@ -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 ) diff --git a/api/services/credit_pool_service.py b/api/services/credit_pool_service.py index 2e013a068ed..bc4d31b58a0 100644 --- a/api/services/credit_pool_service.py +++ b/api/services/credit_pool_service.py @@ -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"] diff --git a/api/services/message_service.py b/api/services/message_service.py index 47efe675d89..cb5ba0306f2 100644 --- a/api/services/message_service.py +++ b/api/services/message_service.py @@ -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, diff --git a/api/services/summary_index_service.py b/api/services/summary_index_service.py index a8506f95cbc..9d820d14659 100644 --- a/api/services/summary_index_service.py +++ b/api/services/summary_index_service.py @@ -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, diff --git a/api/services/vector_service.py b/api/services/vector_service.py index 2ecd0d4f2e8..9d0505ed78e 100644 --- a/api/services/vector_service.py +++ b/api/services/vector_service.py @@ -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 ): diff --git a/api/services/workflow_generator_service.py b/api/services/workflow_generator_service.py index 2ac697d28a9..66c31e60401 100644 --- a/api/services/workflow_generator_service.py +++ b/api/services/workflow_generator_service.py @@ -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, diff --git a/api/tasks/batch_create_segment_to_index_task.py b/api/tasks/batch_create_segment_to_index_task.py index 1fbc4add503..89ee3d17d86 100644 --- a/api/tasks/batch_create_segment_to_index_task.py +++ b/api/tasks/batch_create_segment_to_index_task.py @@ -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, diff --git a/api/tasks/regenerate_summary_index_task.py b/api/tasks/regenerate_summary_index_task.py index 7a4a6331350..69aeba389bf 100644 --- a/api/tasks/regenerate_summary_index_task.py +++ b/api/tasks/regenerate_summary_index_task.py @@ -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", diff --git a/api/tests/unit_tests/core/agent/test_cot_agent_runner.py b/api/tests/unit_tests/core/agent/test_cot_agent_runner.py index 9537d4387c6..771b819fa02 100644 --- a/api/tests/unit_tests/core/agent/test_cot_agent_runner.py +++ b/api/tests/unit_tests/core/agent/test_cot_agent_runner.py @@ -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): diff --git a/api/tests/unit_tests/core/agent/test_fc_agent_runner.py b/api/tests/unit_tests/core/agent/test_fc_agent_runner.py index c8f186ce7f5..4d98d9a7b37 100644 --- a/api/tests/unit_tests/core/agent/test_fc_agent_runner.py +++ b/api/tests/unit_tests/core/agent/test_fc_agent_runner.py @@ -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 diff --git a/api/tests/unit_tests/core/app/apps/chat/test_app_generator_and_runner.py b/api/tests/unit_tests/core/app/apps/chat/test_app_generator_and_runner.py index 905e711414a..5f31cf6b49a 100644 --- a/api/tests/unit_tests/core/app/apps/chat/test_app_generator_and_runner.py +++ b/api/tests/unit_tests/core/app/apps/chat/test_app_generator_and_runner.py @@ -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, diff --git a/api/tests/unit_tests/core/app/apps/completion/test_app_runner.py b/api/tests/unit_tests/core/app/apps/completion/test_app_runner.py index 943d1d01d20..894b5010e76 100644 --- a/api/tests/unit_tests/core/app/apps/completion/test_app_runner.py +++ b/api/tests/unit_tests/core/app/apps/completion/test_app_runner.py @@ -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): diff --git a/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_single_node.py b/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_single_node.py index bdfedd82fe4..3ae4462f9b8 100644 --- a/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_single_node.py +++ b/api/tests/unit_tests/core/app/apps/test_workflow_app_runner_single_node.py @@ -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() diff --git a/api/tests/unit_tests/core/app/entities/test_app_invoke_entities.py b/api/tests/unit_tests/core/app/entities/test_app_invoke_entities.py index 86c80985c45..60d856dd416 100644 --- a/api/tests/unit_tests/core/app/entities/test_app_invoke_entities.py +++ b/api/tests/unit_tests/core/app/entities/test_app_invoke_entities.py @@ -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 diff --git a/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline_core.py b/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline_core.py index 0d9c4b11cd9..7c696955c97 100644 --- a/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline_core.py +++ b/api/tests/unit_tests/core/app/task_pipeline/test_easy_ui_based_generate_task_pipeline_core.py @@ -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()) diff --git a/api/tests/unit_tests/core/app/test_llm_quota.py b/api/tests/unit_tests/core/app/test_llm_quota.py index 55e4293dae9..ffeb2e5e8fd 100644 --- a/api/tests/unit_tests/core/app/test_llm_quota.py +++ b/api/tests/unit_tests/core/app/test_llm_quota.py @@ -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, ) diff --git a/api/tests/unit_tests/core/moderation/test_input_moderation.py b/api/tests/unit_tests/core/moderation/test_input_moderation.py index 2dbc80cf14d..8822b60d68d 100644 --- a/api/tests/unit_tests/core/moderation/test_input_moderation.py +++ b/api/tests/unit_tests/core/moderation/test_input_moderation.py @@ -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 diff --git a/api/tests/unit_tests/core/rag/indexing/test_indexing_runner.py b/api/tests/unit_tests/core/rag/indexing/test_indexing_runner.py index 926fe1baf1e..a131e232369 100644 --- a/api/tests/unit_tests/core/rag/indexing/test_indexing_runner.py +++ b/api/tests/unit_tests/core/rag/indexing/test_indexing_runner.py @@ -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 diff --git a/api/tests/unit_tests/core/test_model_context.py b/api/tests/unit_tests/core/test_model_context.py new file mode 100644 index 00000000000..c3dc92d7efd --- /dev/null +++ b/api/tests/unit_tests/core/test_model_context.py @@ -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} diff --git a/api/tests/unit_tests/core/test_model_manager.py b/api/tests/unit_tests/core/test_model_manager.py index 18c5245e851..ca15d835734 100644 --- a/api/tests/unit_tests/core/test_model_manager.py +++ b/api/tests/unit_tests/core/test_model_manager.py @@ -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() diff --git a/api/tests/unit_tests/core/workflow/test_node_factory.py b/api/tests/unit_tests/core/workflow/test_node_factory.py index f5dd034a649..c92c6a0e588 100644 --- a/api/tests/unit_tests/core/workflow/test_node_factory.py +++ b/api/tests/unit_tests/core/workflow/test_node_factory.py @@ -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 diff --git a/api/tests/unit_tests/core/workflow/test_node_runtime.py b/api/tests/unit_tests/core/workflow/test_node_runtime.py index 3bc5c01ffb3..a8ca5de4ac7 100644 --- a/api/tests/unit_tests/core/workflow/test_node_runtime.py +++ b/api/tests/unit_tests/core/workflow/test_node_runtime.py @@ -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"}, ) diff --git a/api/tests/unit_tests/core/workflow/test_workflow_entry_helpers.py b/api/tests/unit_tests/core/workflow/test_workflow_entry_helpers.py index 520b23faa19..a23a6236487 100644 --- a/api/tests/unit_tests/core/workflow/test_workflow_entry_helpers.py +++ b/api/tests/unit_tests/core/workflow/test_workflow_entry_helpers.py @@ -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="", diff --git a/api/tests/unit_tests/events/test_update_provider_when_message_created.py b/api/tests/unit_tests/events/test_update_provider_when_message_created.py index 248709c6be4..e91e390b5a5 100644 --- a/api/tests/unit_tests/events/test_update_provider_when_message_created.py +++ b/api/tests/unit_tests/events/test_update_provider_when_message_created.py @@ -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", diff --git a/api/tests/unit_tests/services/test_agent_llm_inner_service.py b/api/tests/unit_tests/services/test_agent_llm_inner_service.py index de8fa2fea4f..37eb8f851ad 100644 --- a/api/tests/unit_tests/services/test_agent_llm_inner_service.py +++ b/api/tests/unit_tests/services/test_agent_llm_inner_service.py @@ -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, diff --git a/api/tests/unit_tests/services/test_audio_service.py b/api/tests/unit_tests/services/test_audio_service.py index 1050ef48268..82568916445 100644 --- a/api/tests/unit_tests/services/test_audio_service.py +++ b/api/tests/unit_tests/services/test_audio_service.py @@ -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", diff --git a/api/tests/unit_tests/services/test_credit_pool_service.py b/api/tests/unit_tests/services/test_credit_pool_service.py index 2aeaf5313f0..2781f1de426 100644 --- a/api/tests/unit_tests/services/test_credit_pool_service.py +++ b/api/tests/unit_tests/services/test_credit_pool_service.py @@ -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, + }, ) diff --git a/api/tests/unit_tests/services/test_workflow_generator_service.py b/api/tests/unit_tests/services/test_workflow_generator_service.py index 48b5d9c5523..86f4010bfa2 100644 --- a/api/tests/unit_tests/services/test_workflow_generator_service.py +++ b/api/tests/unit_tests/services/test_workflow_generator_service.py @@ -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