mirror of
https://github.com/langgenius/dify.git
synced 2026-08-29 03:45:08 +08:00
feat(api): track credit usage contexts (#41181)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
@@ -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()
|
||||
|
||||
|
||||
@@ -13,6 +13,8 @@ from core.agent.errors import AgentMaxIterationError
|
||||
from core.app.apps.base_app_queue_manager import PublishFrom
|
||||
from core.app.entities.queue_entities import QueueAgentThoughtEvent, QueueMessageEndEvent, QueueMessageFileEvent
|
||||
from core.app.file_access import grant_upload_file_access
|
||||
from core.credit_usage import CreditUsageAppType, CreditUsageCreatedBy
|
||||
from core.model_context import use_credit_usage_metadata
|
||||
from core.prompt.agent_history_prompt_transform import AgentHistoryPromptTransform
|
||||
from core.tools.entities.tool_entities import ToolInvokeMeta
|
||||
from core.tools.signature import sign_upload_file_preview_url
|
||||
@@ -167,6 +169,12 @@ class FunctionCallAgentRunner(BaseAgentRunner):
|
||||
session.close()
|
||||
|
||||
# invoke model
|
||||
request_metadata: dict[str, object] = {
|
||||
"app_id": self.app_config.app_id,
|
||||
"app_type": CreditUsageAppType.AGENT,
|
||||
"created_by": CreditUsageCreatedBy.APP,
|
||||
}
|
||||
|
||||
chunks: Union[Generator[LLMResultChunk, None, None], LLMResult] = model_instance.invoke_llm(
|
||||
prompt_messages=prompt_messages,
|
||||
model_parameters=app_generate_entity.model_conf.parameters,
|
||||
@@ -174,7 +182,7 @@ class FunctionCallAgentRunner(BaseAgentRunner):
|
||||
stop=app_generate_entity.model_conf.stop,
|
||||
stream=self.stream_tool_call,
|
||||
callbacks=[],
|
||||
request_metadata={"app_id": self.app_config.app_id},
|
||||
request_metadata=request_metadata,
|
||||
)
|
||||
|
||||
tool_calls: list[tuple[str, str, dict[str, Any]]] = []
|
||||
@@ -315,20 +323,21 @@ class FunctionCallAgentRunner(BaseAgentRunner):
|
||||
}
|
||||
else:
|
||||
# invoke tool
|
||||
tool_invoke_response, message_files, tool_invoke_meta = ToolEngine.agent_invoke(
|
||||
session=session,
|
||||
tool=tool_instance,
|
||||
tool_parameters=tool_call_args,
|
||||
user_id=self.user_id,
|
||||
tenant_id=self.tenant_id,
|
||||
message=self.message,
|
||||
invoke_from=self.application_generate_entity.invoke_from,
|
||||
agent_tool_callback=self.agent_callback,
|
||||
trace_manager=trace_manager,
|
||||
app_id=self.application_generate_entity.app_config.app_id,
|
||||
message_id=self.message.id,
|
||||
conversation_id=self.conversation.id,
|
||||
)
|
||||
with use_credit_usage_metadata({"app_type": CreditUsageAppType.AGENT}):
|
||||
tool_invoke_response, message_files, tool_invoke_meta = ToolEngine.agent_invoke(
|
||||
session=session,
|
||||
tool=tool_instance,
|
||||
tool_parameters=tool_call_args,
|
||||
user_id=self.user_id,
|
||||
tenant_id=self.tenant_id,
|
||||
message=self.message,
|
||||
invoke_from=self.application_generate_entity.invoke_from,
|
||||
agent_tool_callback=self.agent_callback,
|
||||
trace_manager=trace_manager,
|
||||
app_id=self.application_generate_entity.app_config.app_id,
|
||||
message_id=self.message.id,
|
||||
conversation_id=self.conversation.id,
|
||||
)
|
||||
session.commit()
|
||||
session.close()
|
||||
# publish files
|
||||
|
||||
@@ -20,6 +20,7 @@ from core.app.entities.app_invoke_entities import (
|
||||
AppGenerateEntity,
|
||||
DifyRunContext,
|
||||
InvokeFrom,
|
||||
get_credit_usage_app_type,
|
||||
)
|
||||
from core.app.entities.queue_entities import (
|
||||
QueueAnnotationReplyEvent,
|
||||
@@ -141,6 +142,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
|
||||
user_id=self.application_generate_entity.user_id,
|
||||
invoke_from=invoke_from,
|
||||
user_from=user_from,
|
||||
app_type=get_credit_usage_app_type(app_config.app_mode),
|
||||
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
|
||||
)
|
||||
elif self.application_generate_entity.single_iteration_run or self.application_generate_entity.single_loop_run:
|
||||
@@ -150,6 +152,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
|
||||
single_iteration_run=self.application_generate_entity.single_iteration_run,
|
||||
single_loop_run=self.application_generate_entity.single_loop_run,
|
||||
user_id=self.application_generate_entity.user_id,
|
||||
app_type=get_credit_usage_app_type(app_config.app_mode),
|
||||
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
|
||||
)
|
||||
else:
|
||||
@@ -220,6 +223,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
|
||||
user_from=user_from,
|
||||
invoke_from=invoke_from,
|
||||
root_node_id=root_node_id,
|
||||
app_type=get_credit_usage_app_type(app_config.app_mode),
|
||||
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
|
||||
)
|
||||
|
||||
@@ -280,6 +284,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
|
||||
user_id=self.application_generate_entity.user_id,
|
||||
user_from=user_from,
|
||||
invoke_from=invoke_from,
|
||||
app_type=get_credit_usage_app_type(app_config.app_mode),
|
||||
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
|
||||
)
|
||||
)
|
||||
|
||||
@@ -18,6 +18,7 @@ from core.app.apps.draft_variable_saver import DraftVariableSaverFactory
|
||||
from core.app.entities.app_invoke_entities import (
|
||||
AdvancedChatAppGenerateEntity,
|
||||
InvokeFrom,
|
||||
get_credit_usage_app_type,
|
||||
)
|
||||
from core.app.entities.queue_entities import (
|
||||
MessageQueueMessage,
|
||||
@@ -371,7 +372,10 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
and features_dict["text_to_speech"].get("autoPlay") == "enabled"
|
||||
):
|
||||
tts_publisher = AppGeneratorTTSPublisher(
|
||||
tenant_id, features_dict["text_to_speech"].get("voice"), features_dict["text_to_speech"].get("language")
|
||||
tenant_id,
|
||||
features_dict["text_to_speech"].get("voice"),
|
||||
features_dict["text_to_speech"].get("language"),
|
||||
get_credit_usage_app_type(self._application_generate_entity.app_config.app_mode),
|
||||
)
|
||||
|
||||
try:
|
||||
|
||||
@@ -45,6 +45,7 @@ from core.app.entities.app_invoke_entities import (
|
||||
InvokeFrom,
|
||||
UserFrom,
|
||||
)
|
||||
from core.credit_usage import CreditUsageAppType
|
||||
from core.db.session_factory import session_factory
|
||||
from core.ops.ops_trace_manager import TraceQueueManager
|
||||
from core.workflow.file_reference import build_file_reference, is_canonical_file_reference
|
||||
@@ -498,6 +499,7 @@ class AgentAppGenerator(MessageBasedAppGenerator):
|
||||
user_id=application_generate_entity.user_id,
|
||||
user_from=user_from,
|
||||
invoke_from=application_generate_entity.invoke_from,
|
||||
app_type=CreditUsageAppType.AGENT_V2,
|
||||
)
|
||||
with session_factory.create_session() as session:
|
||||
agent, config_version, agent_soul = self._resolve_agent_by_id(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -9,6 +9,8 @@ from core.app.apps.base_app_runner import AppRunner
|
||||
from core.app.apps.completion.app_config_manager import CompletionAppConfig
|
||||
from core.app.entities.app_invoke_entities import (
|
||||
CompletionAppGenerateEntity,
|
||||
get_credit_usage_app_type,
|
||||
get_credit_usage_created_by,
|
||||
)
|
||||
from core.callback_handler.index_tool_callback_handler import DatasetIndexToolCallbackHandler
|
||||
from core.db.session_factory import create_session
|
||||
@@ -192,12 +194,16 @@ class CompletionAppRunner(AppRunner):
|
||||
model=application_generate_entity.model_conf.model,
|
||||
)
|
||||
|
||||
request_metadata: dict[str, object] = {"app_id": app_config.app_id}
|
||||
request_metadata["app_type"] = get_credit_usage_app_type(app_config.app_mode)
|
||||
request_metadata["created_by"] = get_credit_usage_created_by(app_config.app_mode)
|
||||
|
||||
invoke_result = model_instance.invoke_llm(
|
||||
prompt_messages=prompt_messages,
|
||||
model_parameters=application_generate_entity.model_conf.parameters,
|
||||
stop=stop,
|
||||
stream=application_generate_entity.stream,
|
||||
request_metadata={"app_id": app_config.app_id},
|
||||
request_metadata=request_metadata,
|
||||
)
|
||||
|
||||
# handle invoke result
|
||||
|
||||
@@ -15,6 +15,7 @@ from core.app.entities.app_invoke_entities import (
|
||||
build_dify_run_context,
|
||||
)
|
||||
from core.app.workflow.layers.persistence import PersistenceWorkflowInfo, WorkflowPersistenceLayer
|
||||
from core.credit_usage import CreditUsageAppType
|
||||
from core.db.session_factory import create_session
|
||||
from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository
|
||||
from core.workflow.node_factory import DifyGraphInitContext, DifyNodeFactory, get_default_root_node_id
|
||||
@@ -302,6 +303,7 @@ class PipelineRunner(WorkflowBasedAppRunner):
|
||||
user_id=self.application_generate_entity.user_id,
|
||||
user_from=user_from,
|
||||
invoke_from=invoke_from,
|
||||
app_type=CreditUsageAppType.RAG_PIPELINE,
|
||||
)
|
||||
graph_init_context = DifyGraphInitContext(
|
||||
workflow_id=workflow.id,
|
||||
|
||||
@@ -13,7 +13,12 @@ from core.app.apps.workflow.command_channels import (
|
||||
)
|
||||
from core.app.apps.workflow.stop_aware_ready_queue import attach_stop_aware_ready_queue
|
||||
from core.app.apps.workflow_app_runner import WorkflowBasedAppRunner
|
||||
from core.app.entities.app_invoke_entities import DifyRunContext, InvokeFrom, WorkflowAppGenerateEntity
|
||||
from core.app.entities.app_invoke_entities import (
|
||||
DifyRunContext,
|
||||
InvokeFrom,
|
||||
WorkflowAppGenerateEntity,
|
||||
get_credit_usage_app_type,
|
||||
)
|
||||
from core.app.workflow.layers.persistence import PersistenceWorkflowInfo, WorkflowPersistenceLayer
|
||||
from core.repositories.factory import WorkflowExecutionRepository, WorkflowNodeExecutionRepository
|
||||
from core.workflow.node_factory import get_default_root_node_id
|
||||
@@ -99,6 +104,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
||||
user_from=user_from,
|
||||
invoke_from=invoke_from,
|
||||
root_node_id=self._root_node_id,
|
||||
app_type=get_credit_usage_app_type(app_config.app_mode),
|
||||
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
|
||||
)
|
||||
elif self.application_generate_entity.single_iteration_run or self.application_generate_entity.single_loop_run:
|
||||
@@ -107,6 +113,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
||||
single_iteration_run=self.application_generate_entity.single_iteration_run,
|
||||
single_loop_run=self.application_generate_entity.single_loop_run,
|
||||
user_id=self.application_generate_entity.user_id,
|
||||
app_type=get_credit_usage_app_type(app_config.app_mode),
|
||||
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
|
||||
)
|
||||
else:
|
||||
@@ -150,6 +157,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
||||
user_from=user_from,
|
||||
invoke_from=invoke_from,
|
||||
root_node_id=root_node_id,
|
||||
app_type=get_credit_usage_app_type(app_config.app_mode),
|
||||
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
|
||||
)
|
||||
|
||||
@@ -210,6 +218,7 @@ class WorkflowAppRunner(WorkflowBasedAppRunner):
|
||||
user_id=self.application_generate_entity.user_id,
|
||||
user_from=user_from,
|
||||
invoke_from=invoke_from,
|
||||
app_type=get_credit_usage_app_type(app_config.app_mode),
|
||||
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
|
||||
)
|
||||
)
|
||||
|
||||
@@ -10,7 +10,7 @@ from core.app.apps.base_app_queue_manager import AppQueueManager
|
||||
from core.app.apps.common.graph_runtime_state_support import GraphRuntimeStateSupport
|
||||
from core.app.apps.common.workflow_response_converter import WorkflowResponseConverter
|
||||
from core.app.apps.draft_variable_saver import DraftVariableSaverFactory
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom, WorkflowAppGenerateEntity
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom, WorkflowAppGenerateEntity, get_credit_usage_app_type
|
||||
from core.app.entities.queue_entities import (
|
||||
AppQueueEvent,
|
||||
MessageQueueMessage,
|
||||
@@ -268,7 +268,10 @@ class WorkflowAppGenerateTaskPipeline(GraphRuntimeStateSupport):
|
||||
and features_dict["text_to_speech"].get("autoPlay") == "enabled"
|
||||
):
|
||||
tts_publisher = AppGeneratorTTSPublisher(
|
||||
tenant_id, features_dict["text_to_speech"].get("voice"), features_dict["text_to_speech"].get("language")
|
||||
tenant_id,
|
||||
features_dict["text_to_speech"].get("voice"),
|
||||
features_dict["text_to_speech"].get("language"),
|
||||
get_credit_usage_app_type(self._application_generate_entity.app_config.app_mode),
|
||||
)
|
||||
|
||||
try:
|
||||
|
||||
@@ -34,6 +34,7 @@ from core.app.entities.queue_entities import (
|
||||
QueueWorkflowStartedEvent,
|
||||
QueueWorkflowSucceededEvent,
|
||||
)
|
||||
from core.credit_usage import CreditUsageAppType
|
||||
from core.rag.entities import RetrievalSourceMetadata
|
||||
from core.repositories.human_input_repository import HumanInputFormSubmissionRepository
|
||||
from core.workflow.node_factory import (
|
||||
@@ -124,6 +125,7 @@ class WorkflowBasedAppRunner:
|
||||
tenant_id: str = "",
|
||||
user_id: str = "",
|
||||
root_node_id: str | None = None,
|
||||
app_type: CreditUsageAppType | None = None,
|
||||
trace_session_id: str | None = None,
|
||||
) -> Graph:
|
||||
"""
|
||||
@@ -145,6 +147,7 @@ class WorkflowBasedAppRunner:
|
||||
user_id=user_id,
|
||||
user_from=user_from,
|
||||
invoke_from=invoke_from,
|
||||
app_type=app_type,
|
||||
trace_session_id=trace_session_id,
|
||||
)
|
||||
graph_init_context = DifyGraphInitContext(
|
||||
@@ -179,6 +182,7 @@ class WorkflowBasedAppRunner:
|
||||
single_loop_run: Any | None = None,
|
||||
*,
|
||||
user_id: str,
|
||||
app_type: CreditUsageAppType | None = None,
|
||||
trace_session_id: str | None = None,
|
||||
) -> tuple[Graph, VariablePool, GraphRuntimeState]:
|
||||
"""
|
||||
@@ -217,6 +221,7 @@ class WorkflowBasedAppRunner:
|
||||
node_type_filter_key="iteration_id",
|
||||
node_type_label="iteration",
|
||||
user_id=user_id,
|
||||
app_type=app_type,
|
||||
trace_session_id=trace_session_id,
|
||||
)
|
||||
elif single_loop_run:
|
||||
@@ -228,6 +233,7 @@ class WorkflowBasedAppRunner:
|
||||
node_type_filter_key="loop_id",
|
||||
node_type_label="loop",
|
||||
user_id=user_id,
|
||||
app_type=app_type,
|
||||
trace_session_id=trace_session_id,
|
||||
)
|
||||
else:
|
||||
@@ -247,6 +253,7 @@ class WorkflowBasedAppRunner:
|
||||
node_type_label: str = "node", # 'iteration' or 'loop' for error messages
|
||||
*,
|
||||
user_id: str = "",
|
||||
app_type: CreditUsageAppType | None = None,
|
||||
trace_session_id: str | None = None,
|
||||
) -> tuple[Graph, VariablePool]:
|
||||
"""
|
||||
@@ -313,6 +320,7 @@ class WorkflowBasedAppRunner:
|
||||
user_id=user_id,
|
||||
user_from=UserFrom.ACCOUNT,
|
||||
invoke_from=InvokeFrom.DEBUGGER,
|
||||
app_type=app_type,
|
||||
trace_session_id=trace_session_id,
|
||||
)
|
||||
graph_init_context = DifyGraphInitContext(
|
||||
|
||||
@@ -6,9 +6,19 @@ from pydantic import BaseModel, ConfigDict, Field, JsonValue, ValidationInfo, fi
|
||||
|
||||
from constants import UUID_NIL
|
||||
from core.app.app_config.entities import EasyUIBasedAppConfig, WorkflowUIBasedAppConfig
|
||||
from core.credit_usage import (
|
||||
CreditUsageAppType,
|
||||
CreditUsageAppTypeInput,
|
||||
CreditUsageCreatedBy,
|
||||
CreditUsageCreatedByInput,
|
||||
created_by_from_app_type,
|
||||
normalize_credit_usage_app_type,
|
||||
normalize_credit_usage_created_by,
|
||||
)
|
||||
from core.entities.provider_configuration import ProviderModelBundle
|
||||
from graphon.file import File, FileUploadConfig
|
||||
from graphon.model_runtime.entities.model_entities import AIModelEntity
|
||||
from models.model import AppMode
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from core.ops.ops_trace_manager import TraceQueueManager
|
||||
@@ -56,14 +66,58 @@ class InvokeFrom(StrEnum):
|
||||
return self in (InvokeFrom.DEBUGGER, InvokeFrom.EXPLORE)
|
||||
|
||||
|
||||
def get_credit_usage_app_type(app_mode: AppMode | str | None) -> CreditUsageAppType:
|
||||
"""Return the top-level application type for an app mode."""
|
||||
if app_mode is None:
|
||||
return CreditUsageAppType.UNKNOWN
|
||||
|
||||
try:
|
||||
normalized_app_mode = app_mode if isinstance(app_mode, AppMode) else AppMode.value_of(str(app_mode))
|
||||
except ValueError:
|
||||
return CreditUsageAppType.UNKNOWN
|
||||
|
||||
app_mode_mapping = {
|
||||
AppMode.CHAT: CreditUsageAppType.CHATBOT,
|
||||
AppMode.ADVANCED_CHAT: CreditUsageAppType.CHATFLOW,
|
||||
AppMode.WORKFLOW: CreditUsageAppType.WORKFLOW,
|
||||
AppMode.AGENT_CHAT: CreditUsageAppType.AGENT,
|
||||
AppMode.AGENT: CreditUsageAppType.AGENT_V2,
|
||||
AppMode.COMPLETION: CreditUsageAppType.COMPLETION,
|
||||
AppMode.CHANNEL: CreditUsageAppType.CHANNEL,
|
||||
AppMode.RAG_PIPELINE: CreditUsageAppType.RAG_PIPELINE,
|
||||
}
|
||||
return app_mode_mapping.get(normalized_app_mode, CreditUsageAppType.UNKNOWN)
|
||||
|
||||
|
||||
def get_credit_usage_created_by(app_mode: AppMode | str | None) -> CreditUsageCreatedBy:
|
||||
"""Return the direct app feature for an app mode."""
|
||||
return created_by_from_app_type(get_credit_usage_app_type(app_mode))
|
||||
|
||||
|
||||
class DifyRunContext(BaseModel):
|
||||
tenant_id: str
|
||||
app_id: str
|
||||
user_id: str
|
||||
user_from: UserFrom
|
||||
invoke_from: InvokeFrom
|
||||
app_type: CreditUsageAppType | None = None
|
||||
created_by: CreditUsageCreatedBy | None = None
|
||||
trace_session_id: str | None = None
|
||||
|
||||
@field_validator("created_by", mode="before")
|
||||
@classmethod
|
||||
def normalize_created_by(cls, value: object) -> CreditUsageCreatedBy | None:
|
||||
if value is None:
|
||||
return None
|
||||
return normalize_credit_usage_created_by(value)
|
||||
|
||||
@field_validator("app_type", mode="before")
|
||||
@classmethod
|
||||
def normalize_app_type(cls, value: object) -> CreditUsageAppType | None:
|
||||
if value is None:
|
||||
return None
|
||||
return normalize_credit_usage_app_type(value)
|
||||
|
||||
|
||||
def build_dify_run_context(
|
||||
*,
|
||||
@@ -72,6 +126,8 @@ def build_dify_run_context(
|
||||
user_id: str,
|
||||
user_from: UserFrom,
|
||||
invoke_from: InvokeFrom,
|
||||
app_type: CreditUsageAppTypeInput = None,
|
||||
created_by: CreditUsageCreatedByInput = None,
|
||||
trace_session_id: str | None = None,
|
||||
extra_context: Mapping[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
@@ -88,6 +144,8 @@ def build_dify_run_context(
|
||||
user_id=user_id,
|
||||
user_from=user_from,
|
||||
invoke_from=invoke_from,
|
||||
app_type=normalize_credit_usage_app_type(app_type) if app_type is not None else None,
|
||||
created_by=normalize_credit_usage_created_by(created_by) if created_by is not None else None,
|
||||
trace_session_id=trace_session_id,
|
||||
)
|
||||
return run_context
|
||||
|
||||
@@ -14,6 +14,14 @@ from sqlalchemy import select
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from configs import dify_config
|
||||
from core.credit_usage import (
|
||||
CreditUsageAppType,
|
||||
CreditUsageAppTypeInput,
|
||||
CreditUsageCreatedBy,
|
||||
CreditUsageCreatedByInput,
|
||||
normalize_credit_usage_app_type,
|
||||
normalize_credit_usage_created_by,
|
||||
)
|
||||
from core.entities.model_entities import ModelStatus
|
||||
from core.entities.provider_entities import ProviderQuotaType, QuotaUnit
|
||||
from core.errors.error import QuotaExceededError
|
||||
@@ -25,7 +33,12 @@ from graphon.model_runtime.entities.model_entities import ModelType
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
from models.provider import Provider, ProviderType
|
||||
from models.provider_ids import ModelProviderID
|
||||
from services.credit_pool_service import CreditPoolReservation, CreditPoolService
|
||||
from services.credit_pool_service import (
|
||||
CREDIT_USAGE_APP_TYPE_META_KEY,
|
||||
CREDIT_USAGE_CREATED_BY_META_KEY,
|
||||
CreditPoolReservation,
|
||||
CreditPoolService,
|
||||
)
|
||||
|
||||
|
||||
class ModelQuotaReservationState(StrEnum):
|
||||
@@ -45,6 +58,8 @@ class ModelQuotaReservation:
|
||||
provider_configuration: Any
|
||||
quota_unit: QuotaUnit | None = None
|
||||
credit_pool_reservation: CreditPoolReservation | None = None
|
||||
app_type: CreditUsageAppType | None = None
|
||||
created_by: CreditUsageCreatedBy | None = None
|
||||
requires_settlement: bool = False
|
||||
_state: ModelQuotaReservationState = field(default=ModelQuotaReservationState.RESERVED, init=False, repr=False)
|
||||
|
||||
@@ -76,6 +91,10 @@ class ModelQuotaReservation:
|
||||
provider=self.provider,
|
||||
provider_configuration=self.provider_configuration,
|
||||
used_quota=used_quota,
|
||||
model_type=self.model_type,
|
||||
model=self.model,
|
||||
app_type=self.app_type,
|
||||
created_by=self.created_by,
|
||||
)
|
||||
|
||||
self._state = ModelQuotaReservationState.COMMITTED
|
||||
@@ -121,15 +140,21 @@ def reserve_model_quota_for_model(
|
||||
model_type: ModelType,
|
||||
model: str,
|
||||
request_id: str | None = None,
|
||||
app_type: CreditUsageAppTypeInput = None,
|
||||
created_by: CreditUsageCreatedByInput = None,
|
||||
) -> ModelQuotaReservation:
|
||||
"""Reserve system-hosted model quota before invoking the provider."""
|
||||
provider_configuration = _get_provider_configuration(tenant_id=tenant_id, provider=provider)
|
||||
effective_app_type = normalize_credit_usage_app_type(app_type)
|
||||
effective_created_by = normalize_credit_usage_created_by(created_by)
|
||||
reservation = ModelQuotaReservation(
|
||||
tenant_id=tenant_id,
|
||||
provider=provider,
|
||||
model_type=model_type,
|
||||
model=model,
|
||||
provider_configuration=provider_configuration,
|
||||
app_type=effective_app_type,
|
||||
created_by=effective_created_by,
|
||||
)
|
||||
if provider_configuration.using_provider_type != ProviderType.SYSTEM:
|
||||
return reservation
|
||||
@@ -166,6 +191,8 @@ def reserve_model_quota_for_model(
|
||||
"model_type": model_type.value,
|
||||
"model": model,
|
||||
}
|
||||
reservation_meta[CREDIT_USAGE_CREATED_BY_META_KEY] = effective_created_by
|
||||
reservation_meta[CREDIT_USAGE_APP_TYPE_META_KEY] = effective_app_type
|
||||
reservation.credit_pool_reservation = CreditPoolService.reserve_credits(
|
||||
tenant_id=tenant_id,
|
||||
credits_required=amount,
|
||||
@@ -183,7 +210,13 @@ def reserve_model_quota_for_model(
|
||||
|
||||
|
||||
def reserve_llm_quota_for_model(
|
||||
*, tenant_id: str, provider: str, model: str, request_id: str | None = None
|
||||
*,
|
||||
tenant_id: str,
|
||||
provider: str,
|
||||
model: str,
|
||||
request_id: str | None = None,
|
||||
app_type: CreditUsageAppTypeInput = None,
|
||||
created_by: CreditUsageCreatedByInput = None,
|
||||
) -> ModelQuotaReservation:
|
||||
"""Reserve system-hosted LLM quota before invoking the provider."""
|
||||
return reserve_model_quota_for_model(
|
||||
@@ -192,6 +225,8 @@ def reserve_llm_quota_for_model(
|
||||
model_type=ModelType.LLM,
|
||||
model=model,
|
||||
request_id=request_id,
|
||||
app_type=app_type,
|
||||
created_by=created_by,
|
||||
)
|
||||
|
||||
|
||||
@@ -286,13 +321,31 @@ def _deduct_free_model_quota(
|
||||
raise QuotaExceededError(f"Model provider {provider} quota exceeded.")
|
||||
|
||||
|
||||
def _deduct_used_model_quota(*, tenant_id: str, provider: str, provider_configuration, used_quota: int | None) -> None:
|
||||
def _deduct_used_model_quota(
|
||||
*,
|
||||
tenant_id: str,
|
||||
provider: str,
|
||||
provider_configuration,
|
||||
used_quota: int | None,
|
||||
model_type: ModelType | None = None,
|
||||
model: str | None = None,
|
||||
app_type: CreditUsageAppTypeInput = None,
|
||||
created_by: CreditUsageCreatedByInput = None,
|
||||
) -> None:
|
||||
"""Apply a resolved model quota charge against the current provider quota bucket."""
|
||||
if provider_configuration.using_provider_type != ProviderType.SYSTEM:
|
||||
return
|
||||
|
||||
system_configuration = provider_configuration.system_configuration
|
||||
if used_quota is not None and system_configuration.current_quota_type is not None:
|
||||
metadata: dict[str, object] = {"provider": provider}
|
||||
if model is not None:
|
||||
metadata["model"] = model
|
||||
if model_type is not None:
|
||||
metadata["model_type"] = model_type.value
|
||||
metadata[CREDIT_USAGE_APP_TYPE_META_KEY] = normalize_credit_usage_app_type(app_type)
|
||||
metadata[CREDIT_USAGE_CREATED_BY_META_KEY] = normalize_credit_usage_created_by(created_by)
|
||||
|
||||
match system_configuration.current_quota_type:
|
||||
case ProviderQuotaType.TRIAL:
|
||||
from services.credit_pool_service import CreditPoolService
|
||||
@@ -300,6 +353,7 @@ def _deduct_used_model_quota(*, tenant_id: str, provider: str, provider_configur
|
||||
CreditPoolService.deduct_credits_capped(
|
||||
tenant_id=tenant_id,
|
||||
credits_required=used_quota,
|
||||
metadata=metadata,
|
||||
session=db.session(),
|
||||
)
|
||||
case ProviderQuotaType.PAID:
|
||||
@@ -309,6 +363,7 @@ def _deduct_used_model_quota(*, tenant_id: str, provider: str, provider_configur
|
||||
tenant_id=tenant_id,
|
||||
credits_required=used_quota,
|
||||
pool_type="paid",
|
||||
metadata=metadata,
|
||||
session=db.session(),
|
||||
)
|
||||
case ProviderQuotaType.FREE:
|
||||
@@ -322,7 +377,15 @@ def _deduct_used_model_quota(*, tenant_id: str, provider: str, provider_configur
|
||||
return
|
||||
|
||||
|
||||
def deduct_llm_quota_for_model(*, tenant_id: str, provider: str, model: str, usage: LLMUsage) -> None:
|
||||
def deduct_llm_quota_for_model(
|
||||
*,
|
||||
tenant_id: str,
|
||||
provider: str,
|
||||
model: str,
|
||||
usage: LLMUsage,
|
||||
app_type: CreditUsageAppTypeInput = None,
|
||||
created_by: CreditUsageCreatedByInput = None,
|
||||
) -> None:
|
||||
"""Deduct tenant-bound quota for the resolved LLM model identity."""
|
||||
provider_configuration = _get_provider_configuration(tenant_id=tenant_id, provider=provider)
|
||||
used_quota = _resolve_llm_used_quota(
|
||||
@@ -335,6 +398,10 @@ def deduct_llm_quota_for_model(*, tenant_id: str, provider: str, model: str, usa
|
||||
provider=provider,
|
||||
provider_configuration=provider_configuration,
|
||||
used_quota=used_quota,
|
||||
model_type=ModelType.LLM,
|
||||
model=model,
|
||||
app_type=app_type,
|
||||
created_by=created_by,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -13,6 +13,7 @@ from core.app.entities.app_invoke_entities import (
|
||||
AgentChatAppGenerateEntity,
|
||||
ChatAppGenerateEntity,
|
||||
CompletionAppGenerateEntity,
|
||||
get_credit_usage_app_type,
|
||||
)
|
||||
from core.app.entities.queue_entities import (
|
||||
QueueAgentMessageEvent,
|
||||
@@ -231,7 +232,10 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline[EasyUIAppGenerat
|
||||
and text_to_speech_dict.get("enabled")
|
||||
):
|
||||
publisher = AppGeneratorTTSPublisher(
|
||||
tenant_id, text_to_speech_dict.get("voice", ""), text_to_speech_dict.get("language", None)
|
||||
tenant_id,
|
||||
text_to_speech_dict.get("voice", ""),
|
||||
text_to_speech_dict.get("language", None),
|
||||
get_credit_usage_app_type(self._app_config.app_mode),
|
||||
)
|
||||
try:
|
||||
for response in self._process_stream_response(publisher=publisher, trace_manager=trace_manager):
|
||||
|
||||
@@ -17,6 +17,7 @@ from core.base.tts.audio_mime import (
|
||||
get_model_audio_mime_type,
|
||||
inspect_audio_stream,
|
||||
)
|
||||
from core.credit_usage import CreditUsageAppType, CreditUsageCreatedBy
|
||||
from core.model_manager import ModelManager
|
||||
from graphon.model_runtime.entities.message_entities import TextPromptMessageContent
|
||||
from graphon.model_runtime.entities.model_entities import ModelType
|
||||
@@ -32,14 +33,25 @@ class AudioTrunk:
|
||||
|
||||
|
||||
class AppGeneratorTTSPublisher:
|
||||
def __init__(self, tenant_id: str, voice: str, language: str | None = None):
|
||||
def __init__(
|
||||
self,
|
||||
tenant_id: str,
|
||||
voice: str,
|
||||
language: str | None = None,
|
||||
app_type: CreditUsageAppType = CreditUsageAppType.UNKNOWN,
|
||||
created_by: CreditUsageCreatedBy = CreditUsageCreatedBy.AUDIO,
|
||||
):
|
||||
self.logger = logging.getLogger(__name__)
|
||||
self.tenant_id = tenant_id
|
||||
self.msg_text = ""
|
||||
self._audio_queue: queue.Queue[AudioTrunk] = queue.Queue()
|
||||
self._msg_queue: queue.Queue[WorkflowQueueMessage | MessageQueueMessage | None] = queue.Queue()
|
||||
self.match = re.compile(r"[。.!?]")
|
||||
self.model_manager = ModelManager.for_tenant(tenant_id=self.tenant_id, user_id="responding_tts")
|
||||
self.model_manager = ModelManager.for_tenant(
|
||||
tenant_id=self.tenant_id,
|
||||
user_id="responding_tts",
|
||||
request_metadata={"app_type": app_type, "created_by": created_by},
|
||||
)
|
||||
self.model_instance = self.model_manager.get_default_model_instance(
|
||||
tenant_id=self.tenant_id, model_type=ModelType.TTS
|
||||
)
|
||||
|
||||
@@ -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
|
||||
@@ -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,
|
||||
|
||||
@@ -9,6 +9,7 @@ from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from core.app.app_config.entities import ModelConfig
|
||||
from core.app.entities.app_invoke_entities import CreditUsageCreatedBy
|
||||
from core.llm_generator.entities import RuleCodeGeneratePayload, RuleGeneratePayload, RuleStructuredOutputPayload
|
||||
from core.llm_generator.output_parser.rule_config_generator import RuleConfigGeneratorOutputParser
|
||||
from core.llm_generator.output_parser.suggested_questions_after_answer import SuggestedQuestionsAfterAnswerOutputParser
|
||||
@@ -22,6 +23,7 @@ from core.llm_generator.prompts import (
|
||||
SYSTEM_STRUCTURED_OUTPUT_GENERATE,
|
||||
WORKFLOW_RULE_CONFIG_PROMPT_GENERATE_TEMPLATE,
|
||||
)
|
||||
from core.model_context import with_credit_usage_created_by
|
||||
from core.model_manager import ModelManager
|
||||
from core.ops.entities.trace_entity import TraceTaskName
|
||||
from core.ops.ops_trace_manager import TraceQueueManager, TraceTask
|
||||
@@ -187,6 +189,7 @@ class LLMGenerator:
|
||||
logger.debug("Failed to emit prompt_generation telemetry", exc_info=True)
|
||||
|
||||
@classmethod
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.CONVERSATION_NAME)
|
||||
def generate_conversation_name(
|
||||
cls,
|
||||
tenant_id: str,
|
||||
@@ -255,6 +258,7 @@ class LLMGenerator:
|
||||
return name
|
||||
|
||||
@classmethod
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.SUGGESTED_QUESTIONS)
|
||||
def generate_suggested_questions_after_answer(
|
||||
cls,
|
||||
tenant_id: str,
|
||||
@@ -342,6 +346,7 @@ class LLMGenerator:
|
||||
return questions
|
||||
|
||||
@classmethod
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.WORKFLOW_INSTRUCTION_SUGGESTIONS)
|
||||
def generate_workflow_instruction_suggestions(
|
||||
cls,
|
||||
tenant_id: str,
|
||||
@@ -457,6 +462,7 @@ class LLMGenerator:
|
||||
return "\n\n".join(sections) + "\n\n"
|
||||
|
||||
@classmethod
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.RULE_CONFIG)
|
||||
def generate_rule_config(cls, tenant_id: str, args: RuleGeneratePayload, *, app_id: str | None = None):
|
||||
output_parser = RuleConfigGeneratorOutputParser()
|
||||
|
||||
@@ -624,6 +630,7 @@ class LLMGenerator:
|
||||
return rule_config
|
||||
|
||||
@classmethod
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.CODE_GENERATION)
|
||||
def generate_code(
|
||||
cls,
|
||||
tenant_id: str,
|
||||
@@ -692,6 +699,7 @@ class LLMGenerator:
|
||||
return {"code": generated_code, "language": args.code_language, "error": ""}
|
||||
|
||||
@classmethod
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.QA_DOCUMENT)
|
||||
def generate_qa_document(cls, tenant_id: str, query, document_language: str):
|
||||
prompt = GENERATOR_QA_PROMPT.format(language=document_language)
|
||||
|
||||
@@ -719,6 +727,7 @@ class LLMGenerator:
|
||||
return answer.strip()
|
||||
|
||||
@classmethod
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.STRUCTURED_OUTPUT)
|
||||
def generate_structured_output(
|
||||
cls, tenant_id: str, args: RuleStructuredOutputPayload, *, app_id: str | None = None
|
||||
) -> StructuredOutputResultDict:
|
||||
@@ -776,6 +785,7 @@ class LLMGenerator:
|
||||
return {"output": generated_output, "error": ""}
|
||||
|
||||
@staticmethod
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.INSTRUCTION_MODIFICATION)
|
||||
def instruction_modify_legacy(
|
||||
tenant_id: str,
|
||||
flow_id: str,
|
||||
@@ -821,6 +831,7 @@ class LLMGenerator:
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.INSTRUCTION_MODIFICATION)
|
||||
def instruction_modify_workflow(
|
||||
tenant_id: str,
|
||||
flow_id: str,
|
||||
|
||||
@@ -56,6 +56,7 @@ def invoke_llm_with_structured_output(
|
||||
stop: list[str] | None = None,
|
||||
stream: Literal[True],
|
||||
callbacks: list[Callback] | None = None,
|
||||
request_metadata: Mapping[str, object] | None = None,
|
||||
) -> Generator[LLMResultChunkWithStructuredOutput, None, None]: ...
|
||||
@overload
|
||||
def invoke_llm_with_structured_output(
|
||||
@@ -70,6 +71,7 @@ def invoke_llm_with_structured_output(
|
||||
stop: list[str] | None = None,
|
||||
stream: Literal[False],
|
||||
callbacks: list[Callback] | None = None,
|
||||
request_metadata: Mapping[str, object] | None = None,
|
||||
) -> LLMResultWithStructuredOutput: ...
|
||||
@overload
|
||||
def invoke_llm_with_structured_output(
|
||||
@@ -84,6 +86,7 @@ def invoke_llm_with_structured_output(
|
||||
stop: list[str] | None = None,
|
||||
stream: bool = True,
|
||||
callbacks: list[Callback] | None = None,
|
||||
request_metadata: Mapping[str, object] | None = None,
|
||||
) -> LLMResultWithStructuredOutput | Generator[LLMResultChunkWithStructuredOutput, None, None]: ...
|
||||
def invoke_llm_with_structured_output(
|
||||
*,
|
||||
@@ -97,6 +100,7 @@ def invoke_llm_with_structured_output(
|
||||
stop: list[str] | None = None,
|
||||
stream: bool = True,
|
||||
callbacks: list[Callback] | None = None,
|
||||
request_metadata: Mapping[str, object] | None = None,
|
||||
) -> LLMResultWithStructuredOutput | Generator[LLMResultChunkWithStructuredOutput, None, None]:
|
||||
"""
|
||||
Invoke large language model with structured output
|
||||
@@ -139,6 +143,7 @@ def invoke_llm_with_structured_output(
|
||||
stop=stop,
|
||||
stream=stream,
|
||||
callbacks=callbacks,
|
||||
request_metadata=request_metadata,
|
||||
)
|
||||
|
||||
if isinstance(llm_result, LLMResult):
|
||||
|
||||
@@ -0,0 +1,61 @@
|
||||
from collections.abc import Callable, Generator, Mapping
|
||||
from contextlib import contextmanager
|
||||
from contextvars import ContextVar
|
||||
from functools import wraps
|
||||
from typing import Any, Protocol, cast
|
||||
|
||||
from core.credit_usage import CreditUsageCreatedByInput, normalize_credit_usage_created_by
|
||||
|
||||
_credit_usage_metadata: ContextVar[dict[str, object] | None] = ContextVar(
|
||||
"credit_usage_metadata",
|
||||
default=None,
|
||||
)
|
||||
|
||||
|
||||
class _CreditUsageMetadataCarrier(Protocol):
|
||||
_request_metadata: Mapping[str, object] | None
|
||||
|
||||
|
||||
def get_credit_usage_metadata() -> Mapping[str, object] | None:
|
||||
return _credit_usage_metadata.get()
|
||||
|
||||
|
||||
@contextmanager
|
||||
def use_credit_usage_metadata(metadata: Mapping[str, object] | None) -> Generator[None, None, None]:
|
||||
if metadata is None:
|
||||
yield
|
||||
return
|
||||
|
||||
current_metadata = _credit_usage_metadata.get()
|
||||
effective_metadata = {**metadata, **(current_metadata or {})}
|
||||
token = _credit_usage_metadata.set(effective_metadata)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
_credit_usage_metadata.reset(token)
|
||||
|
||||
|
||||
def with_credit_usage_metadata(method: Callable[..., Any]) -> Callable[..., Any]:
|
||||
@wraps(method)
|
||||
def wrapper(self, *args: object, **kwargs: object):
|
||||
carrier = cast(_CreditUsageMetadataCarrier, self)
|
||||
with use_credit_usage_metadata(carrier._request_metadata):
|
||||
return method(self, *args, **kwargs)
|
||||
|
||||
return wrapper
|
||||
|
||||
|
||||
def with_credit_usage_created_by(
|
||||
created_by: CreditUsageCreatedByInput,
|
||||
) -> Callable[[Callable[..., Any]], Callable[..., Any]]:
|
||||
normalized_created_by = normalize_credit_usage_created_by(created_by)
|
||||
|
||||
def decorator(method: Callable[..., Any]) -> Callable[..., Any]:
|
||||
@wraps(method)
|
||||
def wrapper(*args: object, **kwargs: object):
|
||||
with use_credit_usage_metadata({"created_by": normalized_created_by}):
|
||||
return method(*args, **kwargs)
|
||||
|
||||
return wrapper
|
||||
|
||||
return decorator
|
||||
+185
-27
@@ -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
|
||||
|
||||
@@ -3,6 +3,8 @@ from collections.abc import Mapping
|
||||
from typing import Any
|
||||
|
||||
from core.app.app_config.entities import AppConfig
|
||||
from core.app.entities.app_invoke_entities import get_credit_usage_app_type
|
||||
from core.model_context import use_credit_usage_metadata
|
||||
from core.moderation.base import ModerationAction, ModerationError
|
||||
from core.moderation.factory import ModerationFactory
|
||||
from core.ops.entities.trace_entity import TraceTaskName
|
||||
@@ -41,12 +43,13 @@ class InputModeration:
|
||||
sensitive_word_avoidance_config = app_config.sensitive_word_avoidance
|
||||
moderation_type = sensitive_word_avoidance_config.type
|
||||
|
||||
moderation_factory = ModerationFactory(
|
||||
name=moderation_type, app_id=app_id, tenant_id=tenant_id, config=sensitive_word_avoidance_config.config
|
||||
)
|
||||
with use_credit_usage_metadata({"app_type": get_credit_usage_app_type(app_config.app_mode)}):
|
||||
moderation_factory = ModerationFactory(
|
||||
name=moderation_type, app_id=app_id, tenant_id=tenant_id, config=sensitive_word_avoidance_config.config
|
||||
)
|
||||
|
||||
with measure_time() as timer:
|
||||
moderation_result = moderation_factory.moderation_for_inputs(inputs, query)
|
||||
with measure_time() as timer:
|
||||
moderation_result = moderation_factory.moderation_for_inputs(inputs, query)
|
||||
|
||||
if trace_manager:
|
||||
trace_manager.add_trace_task(
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
from typing import Any, override
|
||||
|
||||
from core.credit_usage import CreditUsageCreatedBy
|
||||
from core.model_context import with_credit_usage_created_by
|
||||
from core.model_manager import ModelManager
|
||||
from core.moderation.base import Moderation, ModerationAction, ModerationInputsResult, ModerationOutputsResult
|
||||
from graphon.model_runtime.entities.model_entities import ModelType
|
||||
@@ -53,6 +55,7 @@ class OpenAIModeration(Moderation):
|
||||
flagged=flagged, action=ModerationAction.DIRECT_OUTPUT, preset_response=preset_response
|
||||
)
|
||||
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.MODERATION)
|
||||
def _is_violated(self, inputs: dict[str, Any]):
|
||||
text = "\n".join(str(inputs.values()))
|
||||
model_manager = ModelManager.for_tenant(tenant_id=self.tenant_id)
|
||||
|
||||
@@ -4,7 +4,9 @@ from collections.abc import Generator
|
||||
from typing import Any
|
||||
|
||||
from core.base.tts.audio_mime import get_model_audio_mime_type, inspect_audio_stream
|
||||
from core.credit_usage import CreditUsageCreatedBy
|
||||
from core.llm_generator.output_parser.structured_output import invoke_llm_with_structured_output
|
||||
from core.model_context import with_credit_usage_created_by
|
||||
from core.model_manager import ModelManager
|
||||
from core.plugin.backwards_invocation.base import BaseBackwardsInvocation
|
||||
from core.plugin.entities.request import (
|
||||
@@ -37,6 +39,7 @@ from models.account import Tenant
|
||||
|
||||
class PluginModelBackwardsInvocation(BaseBackwardsInvocation):
|
||||
@staticmethod
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.PLUGIN_API)
|
||||
def _get_bound_model_instance(
|
||||
*,
|
||||
tenant_id: str,
|
||||
|
||||
@@ -9,9 +9,11 @@ from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from core.app.file_access import DatabaseFileAccessController
|
||||
from core.credit_usage import CreditUsageCreatedBy
|
||||
from core.db.session_factory import session_factory
|
||||
from core.entities.knowledge_entities import PreviewDetail
|
||||
from core.llm_generator.prompts import DEFAULT_GENERATOR_SUMMARY_PROMPT
|
||||
from core.model_context import with_credit_usage_created_by
|
||||
from core.model_manager import ModelManager
|
||||
from core.rag.cleaner.clean_processor import CleanProcessor
|
||||
from core.rag.datasource.keyword.keyword_factory import Keyword
|
||||
@@ -376,6 +378,7 @@ class ParagraphIndexProcessor(BaseIndexProcessor):
|
||||
return preview_texts
|
||||
|
||||
@staticmethod
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_INDEXING)
|
||||
def generate_summary(
|
||||
tenant_id: str,
|
||||
text: str,
|
||||
|
||||
@@ -3,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,
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -19,13 +19,20 @@ from core.app.app_config.entities import (
|
||||
MetadataFilteringCondition,
|
||||
ModelConfig,
|
||||
)
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom, ModelConfigWithCredentialsEntity
|
||||
from core.app.entities.app_invoke_entities import (
|
||||
CreditUsageCreatedBy,
|
||||
EasyUIBasedAppGenerateEntity,
|
||||
InvokeFrom,
|
||||
ModelConfigWithCredentialsEntity,
|
||||
get_credit_usage_app_type,
|
||||
)
|
||||
from core.app.file_access import grant_retriever_segment_access, grant_upload_file_access
|
||||
from core.callback_handler.index_tool_callback_handler import DatasetIndexToolCallbackHandler
|
||||
from core.db.session_factory import session_factory
|
||||
from core.entities.agent_entities import PlanningStrategy
|
||||
from core.entities.model_entities import ModelStatus
|
||||
from core.memory.token_buffer_memory import TokenBufferMemory
|
||||
from core.model_context import with_credit_usage_created_by, with_credit_usage_metadata
|
||||
from core.model_manager import ModelInstance, ModelManager
|
||||
from core.ops.entities.trace_entity import TraceTaskName
|
||||
from core.ops.ops_trace_manager import TraceQueueManager, TraceTask
|
||||
@@ -102,9 +109,19 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class DatasetRetrieval:
|
||||
def __init__(self, application_generate_entity=None):
|
||||
def __init__(self, application_generate_entity: EasyUIBasedAppGenerateEntity | None = None):
|
||||
self.application_generate_entity = application_generate_entity
|
||||
self._llm_usage = LLMUsage.empty_usage()
|
||||
self._request_metadata: dict[str, object] | None = None
|
||||
if application_generate_entity is not None:
|
||||
app_config = application_generate_entity.app_config
|
||||
self._request_metadata = {
|
||||
"app_type": get_credit_usage_app_type(app_config.app_mode),
|
||||
"app_id": app_config.app_id,
|
||||
}
|
||||
|
||||
def set_request_metadata(self, request_metadata: Mapping[str, object] | None) -> None:
|
||||
self._request_metadata = dict(request_metadata) if request_metadata else None
|
||||
|
||||
@property
|
||||
def llm_usage(self) -> LLMUsage:
|
||||
@@ -119,6 +136,8 @@ class DatasetRetrieval:
|
||||
self._llm_usage = self._llm_usage.plus(usage)
|
||||
|
||||
@trace_span()
|
||||
@with_credit_usage_metadata
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_RETRIEVAL)
|
||||
def knowledge_retrieval(self, session: Session, request: KnowledgeRetrievalRequest) -> list[Source]:
|
||||
self._check_knowledge_rate_limit(request.tenant_id)
|
||||
available_datasets = self._get_available_datasets(request.tenant_id, request.dataset_ids)
|
||||
@@ -354,6 +373,8 @@ class DatasetRetrieval:
|
||||
item.metadata.position = position # type: ignore[index]
|
||||
return retrieval_resource_list
|
||||
|
||||
@with_credit_usage_metadata
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_RETRIEVAL)
|
||||
def retrieve(
|
||||
self,
|
||||
session: Session,
|
||||
@@ -1447,6 +1468,7 @@ class DatasetRetrieval:
|
||||
)
|
||||
return filter_documents[:top_k] if top_k else filter_documents
|
||||
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_RETRIEVAL)
|
||||
def get_metadata_filter_condition(
|
||||
self,
|
||||
session: Session,
|
||||
|
||||
@@ -1,12 +1,15 @@
|
||||
from typing import Union
|
||||
|
||||
from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity
|
||||
from core.credit_usage import CreditUsageCreatedBy
|
||||
from core.model_context import with_credit_usage_created_by
|
||||
from core.model_manager import ModelInstance
|
||||
from graphon.model_runtime.entities.llm_entities import LLMResult, LLMUsage
|
||||
from graphon.model_runtime.entities.message_entities import PromptMessageTool, SystemPromptMessage, UserPromptMessage
|
||||
|
||||
|
||||
class FunctionCallMultiDatasetRouter:
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_RETRIEVAL)
|
||||
def invoke(
|
||||
self,
|
||||
query: str,
|
||||
|
||||
@@ -2,6 +2,8 @@ from collections.abc import Generator, Sequence
|
||||
from typing import Any, Union
|
||||
|
||||
from core.app.entities.app_invoke_entities import ModelConfigWithCredentialsEntity
|
||||
from core.credit_usage import CreditUsageCreatedBy
|
||||
from core.model_context import with_credit_usage_created_by
|
||||
from core.model_manager import ModelInstance, ModelManager
|
||||
from core.prompt.advanced_prompt_transform import AdvancedPromptTransform
|
||||
from core.prompt.entities.advanced_prompt_entities import ChatModelMessage, CompletionModelPromptTemplate
|
||||
@@ -135,6 +137,7 @@ class ReactMultiDatasetRouter:
|
||||
return react_decision.tool, usage
|
||||
return None, usage
|
||||
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_RETRIEVAL)
|
||||
def _invoke_llm(
|
||||
self,
|
||||
completion_param: dict[str, Any],
|
||||
|
||||
@@ -6,6 +6,7 @@ from sqlalchemy import select
|
||||
from core.db.session_factory import session_factory
|
||||
from core.rag.index_processor.constant.index_type import IndexTechniqueType
|
||||
from core.rag.index_processor.index_processor_base import SummaryIndexSettingDict
|
||||
from extensions.otel import propagate_context
|
||||
from models.dataset import Dataset, Document, DocumentSegment, DocumentSegmentSummary
|
||||
from services.summary_index_service import SummaryIndexService
|
||||
from tasks.generate_summary_index_task import generate_summary_index_task
|
||||
@@ -92,7 +93,10 @@ class SummaryIndex:
|
||||
# Continue processing other segments
|
||||
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor:
|
||||
futures = [executor.submit(process_segment, segment_id) for segment_id in pending_segment_ids]
|
||||
futures = [
|
||||
executor.submit(propagate_context(process_segment), segment_id)
|
||||
for segment_id in pending_segment_ids
|
||||
]
|
||||
concurrent.futures.wait(futures)
|
||||
else:
|
||||
generate_summary_index_task.delay(dataset_id, document_id, None)
|
||||
|
||||
@@ -4,6 +4,8 @@ from typing import Any, override
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from core.credit_usage import CreditUsageCreatedBy
|
||||
from core.model_context import with_credit_usage_created_by
|
||||
from core.model_manager import ModelManager
|
||||
from core.plugin.entities.parameters import PluginParameterOption
|
||||
from core.tools.builtin_tool.tool import BuiltinTool
|
||||
@@ -17,6 +19,7 @@ from services.model_provider_service import ModelProviderService
|
||||
|
||||
class ASRTool(BuiltinTool):
|
||||
@override
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.AUDIO)
|
||||
def _invoke(
|
||||
self,
|
||||
session: Session,
|
||||
|
||||
@@ -4,6 +4,8 @@ from typing import Any, override
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from core.base.tts.audio_mime import get_model_audio_mime_type, inspect_audio_stream
|
||||
from core.credit_usage import CreditUsageCreatedBy
|
||||
from core.model_context import with_credit_usage_created_by
|
||||
from core.model_manager import ModelManager
|
||||
from core.plugin.entities.parameters import PluginParameterOption
|
||||
from core.tools.builtin_tool.tool import BuiltinTool
|
||||
@@ -15,6 +17,7 @@ from services.model_provider_service import ModelProviderService
|
||||
|
||||
class TTSTool(BuiltinTool):
|
||||
@override
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.AUDIO)
|
||||
def _invoke(
|
||||
self,
|
||||
session: Session,
|
||||
|
||||
@@ -8,6 +8,8 @@ import json
|
||||
from decimal import Decimal
|
||||
from typing import cast
|
||||
|
||||
from core.credit_usage import CreditUsageCreatedBy
|
||||
from core.model_context import with_credit_usage_created_by
|
||||
from core.model_manager import ModelManager
|
||||
from core.tools.entities.tool_entities import ToolProviderType
|
||||
from extensions.ext_database import db
|
||||
@@ -79,6 +81,7 @@ class ModelInvocationUtils:
|
||||
return tokens
|
||||
|
||||
@staticmethod
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.TOOL)
|
||||
def invoke(
|
||||
user_id: str,
|
||||
tenant_id: str,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -20,6 +20,7 @@ from core.db.session_factory import session_factory
|
||||
from core.helper.trace_id_helper import ParentTraceContext
|
||||
from core.llm_generator.output_parser.errors import OutputParserError
|
||||
from core.llm_generator.output_parser.structured_output import invoke_llm_with_structured_output
|
||||
from core.model_context import use_credit_usage_metadata
|
||||
from core.model_manager import ModelInstance, QuotaManagedModelInstance
|
||||
from core.plugin.impl.exc import PluginDaemonClientSideError, PluginInvokeError
|
||||
from core.plugin.impl.plugin import PluginInstaller
|
||||
@@ -303,6 +304,7 @@ class DifyPreparedLLM(LLMProtocol):
|
||||
model_parameters=model_parameters,
|
||||
stop=list(stop or []),
|
||||
stream=stream,
|
||||
request_metadata=self._request_metadata,
|
||||
)
|
||||
|
||||
@override
|
||||
@@ -362,7 +364,7 @@ class DifyPreparedPollingLLM(DifyPreparedLLM, LLMPollingCapableProtocol):
|
||||
self.finalize_llm_polling()
|
||||
|
||||
if isinstance(self._model_instance, QuotaManagedModelInstance):
|
||||
self._polling_quota_reservation = self._model_instance.reserve_quota()
|
||||
self._polling_quota_reservation = self._model_instance._reserve_quota_for_request(self._request_metadata)
|
||||
|
||||
try:
|
||||
polling_result = self._polling_runtime.start_llm_polling(
|
||||
@@ -627,25 +629,29 @@ class DifyToolNodeRuntime(ToolNodeRuntimeProtocol):
|
||||
tool.clear_trace_session_id()
|
||||
|
||||
try:
|
||||
session_maker = self._session_maker or session_factory.get_session_maker()
|
||||
with session_maker.begin() as session:
|
||||
messages = ToolEngine.generic_invoke(
|
||||
session=session,
|
||||
tool=tool,
|
||||
tool_parameters=dict(tool_parameters),
|
||||
user_id=self._run_context.user_id,
|
||||
workflow_tool_callback=callback,
|
||||
workflow_call_depth=workflow_call_depth,
|
||||
app_id=self._run_context.app_id,
|
||||
conversation_id=runtime_binding.conversation_id,
|
||||
)
|
||||
transformed_messages = ToolFileMessageTransformer.transform_tool_invoke_messages(
|
||||
messages=messages,
|
||||
user_id=self._run_context.user_id,
|
||||
tenant_id=self._run_context.tenant_id,
|
||||
conversation_id=runtime_binding.conversation_id,
|
||||
)
|
||||
yield from self._adapt_messages(transformed_messages, provider_name=provider_name)
|
||||
request_metadata = (
|
||||
{"app_type": self._run_context.app_type} if self._run_context.app_type is not None else None
|
||||
)
|
||||
with use_credit_usage_metadata(request_metadata):
|
||||
session_maker = self._session_maker or session_factory.get_session_maker()
|
||||
with session_maker.begin() as session:
|
||||
messages = ToolEngine.generic_invoke(
|
||||
session=session,
|
||||
tool=tool,
|
||||
tool_parameters=dict(tool_parameters),
|
||||
user_id=self._run_context.user_id,
|
||||
workflow_tool_callback=callback,
|
||||
workflow_call_depth=workflow_call_depth,
|
||||
app_id=self._run_context.app_id,
|
||||
conversation_id=runtime_binding.conversation_id,
|
||||
)
|
||||
transformed_messages = ToolFileMessageTransformer.transform_tool_invoke_messages(
|
||||
messages=messages,
|
||||
user_id=self._run_context.user_id,
|
||||
tenant_id=self._run_context.tenant_id,
|
||||
conversation_id=runtime_binding.conversation_id,
|
||||
)
|
||||
yield from self._adapt_messages(transformed_messages, provider_name=provider_name)
|
||||
except Exception as exc:
|
||||
raise self._map_invocation_exception(exc, provider_name=provider_name) from exc
|
||||
|
||||
|
||||
@@ -4,7 +4,9 @@ from typing import TYPE_CHECKING, Any, override
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from core.app.entities.app_invoke_entities import DIFY_RUN_CONTEXT_KEY, DifyRunContext
|
||||
from core.db.session_factory import session_factory
|
||||
from core.model_context import use_credit_usage_metadata
|
||||
from core.rag.index_processor.index_processor import IndexProcessor
|
||||
from core.rag.index_processor.index_processor_base import SummaryIndexSettingDict
|
||||
from core.rag.summary_index.summary_index import SummaryIndex
|
||||
@@ -150,14 +152,17 @@ class KnowledgeIndexNode(Node[KnowledgeIndexNodeData]):
|
||||
):
|
||||
if not document_id:
|
||||
raise KnowledgeIndexNodeError("document_id is required.")
|
||||
rst = self.index_processor.index_and_clean(
|
||||
dataset_id, document_id, original_document_id, chunks, batch, summary_index_setting, session=session
|
||||
)
|
||||
# Summary generation opens independent sessions and must see the indexed rows.
|
||||
session.commit()
|
||||
self.summary_index_service.generate_and_vectorize_summary(
|
||||
dataset_id, document_id, is_preview, summary_index_setting
|
||||
)
|
||||
dify_ctx = DifyRunContext.model_validate(self.require_run_context_value(DIFY_RUN_CONTEXT_KEY))
|
||||
request_metadata = {"app_type": dify_ctx.app_type} if dify_ctx.app_type is not None else None
|
||||
with use_credit_usage_metadata(request_metadata):
|
||||
rst = self.index_processor.index_and_clean(
|
||||
dataset_id, document_id, original_document_id, chunks, batch, summary_index_setting, session=session
|
||||
)
|
||||
# Summary generation opens independent sessions and must see the indexed rows.
|
||||
session.commit()
|
||||
self.summary_index_service.generate_and_vectorize_summary(
|
||||
dataset_id, document_id, is_preview, summary_index_setting
|
||||
)
|
||||
return rst
|
||||
|
||||
@classmethod
|
||||
|
||||
@@ -12,6 +12,7 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from core.app.app_config.entities import DatasetRetrieveConfigEntity
|
||||
from core.app.entities.app_invoke_entities import DIFY_RUN_CONTEXT_KEY, DifyRunContext
|
||||
from core.credit_usage import CreditUsageAppType
|
||||
from core.db.session_factory import session_factory
|
||||
from core.rag.data_post_processor.data_post_processor import RerankingModelDict, WeightsDict
|
||||
from core.rag.retrieval.dataset_retrieval import DatasetRetrieval
|
||||
@@ -185,6 +186,12 @@ class KnowledgeRetrievalNode(Node[KnowledgeRetrievalNodeData]):
|
||||
self, session: Session, node_data: KnowledgeRetrievalNodeData, variables: dict[str, Any]
|
||||
) -> tuple[list[Source], LLMUsage]:
|
||||
dify_ctx = DifyRunContext.model_validate(self.require_run_context_value(DIFY_RUN_CONTEXT_KEY))
|
||||
self._rag_retrieval.set_request_metadata(
|
||||
{
|
||||
"app_id": dify_ctx.app_id,
|
||||
"app_type": dify_ctx.app_type or CreditUsageAppType.UNKNOWN,
|
||||
}
|
||||
)
|
||||
dataset_ids = node_data.dataset_ids
|
||||
query = variables.get("query")
|
||||
attachments = variables.get("attachments")
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
from collections.abc import Mapping
|
||||
from typing import Any, Literal, Protocol
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
@@ -86,4 +87,6 @@ class RAGRetrievalProtocol(Protocol):
|
||||
@property
|
||||
def llm_usage(self) -> LLMUsage: ...
|
||||
|
||||
def set_request_metadata(self, request_metadata: Mapping[str, object] | None) -> None: ...
|
||||
|
||||
def knowledge_retrieval(self, request: KnowledgeRetrievalRequest) -> list[Source]: ...
|
||||
|
||||
@@ -7,9 +7,14 @@ from uuid import uuid4
|
||||
from configs import dify_config
|
||||
from context import capture_current_context
|
||||
from core.app.apps.exc import GenerateTaskStoppedError
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom, build_dify_run_context
|
||||
from core.app.entities.app_invoke_entities import (
|
||||
InvokeFrom,
|
||||
UserFrom,
|
||||
build_dify_run_context,
|
||||
)
|
||||
from core.app.file_access import DatabaseFileAccessController
|
||||
from core.app.workflow.layers.observability import ObservabilityLayer
|
||||
from core.credit_usage import CreditUsageAppType
|
||||
from core.workflow.node_factory import (
|
||||
DifyGraphInitContext,
|
||||
DifyNodeFactory,
|
||||
@@ -228,6 +233,7 @@ class WorkflowEntry:
|
||||
user_id=user_id,
|
||||
user_from=UserFrom.ACCOUNT,
|
||||
invoke_from=InvokeFrom.DEBUGGER,
|
||||
app_type=CreditUsageAppType.WORKFLOW,
|
||||
)
|
||||
graph_init_context = DifyGraphInitContext(
|
||||
workflow_id=workflow.id,
|
||||
@@ -387,6 +393,7 @@ class WorkflowEntry:
|
||||
user_id=user_id,
|
||||
user_from=UserFrom.ACCOUNT,
|
||||
invoke_from=InvokeFrom.DEBUGGER,
|
||||
app_type=CreditUsageAppType.WORKFLOW,
|
||||
)
|
||||
graph_init_context = DifyGraphInitContext(
|
||||
workflow_id="",
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
import logging
|
||||
import time as time_module
|
||||
from collections.abc import Mapping
|
||||
from datetime import datetime
|
||||
from typing import Any, cast
|
||||
from typing import Any, TypedDict, cast
|
||||
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy import update
|
||||
@@ -13,11 +14,14 @@ from core.app.entities.app_invoke_entities import (
|
||||
AgentAppGenerateEntity,
|
||||
AgentChatAppGenerateEntity,
|
||||
ChatAppGenerateEntity,
|
||||
get_credit_usage_app_type,
|
||||
get_credit_usage_created_by,
|
||||
)
|
||||
from core.entities.provider_entities import ProviderQuotaType, QuotaUnit, SystemConfiguration
|
||||
from events.message_event import message_was_created
|
||||
from extensions.ext_database import db
|
||||
from extensions.ext_redis import redis_client, redis_fallback
|
||||
from graphon.model_runtime.entities.model_entities import ModelType
|
||||
from libs import datetime_utils
|
||||
from models.model import Message
|
||||
from models.provider import Provider, ProviderType
|
||||
@@ -83,6 +87,11 @@ class _ProviderUpdateOperation(BaseModel):
|
||||
description: str = "unknown"
|
||||
|
||||
|
||||
class _CreditDeductionContext(TypedDict, total=False):
|
||||
request_id: str | None
|
||||
metadata: Mapping[str, object]
|
||||
|
||||
|
||||
@message_was_created.connect
|
||||
def handle(sender: Message, **kwargs):
|
||||
"""
|
||||
@@ -124,7 +133,18 @@ def handle(sender: Message, **kwargs):
|
||||
model_config = application_generate_entity.model_conf
|
||||
provider_model_bundle = model_config.provider_model_bundle
|
||||
provider_configuration = provider_model_bundle.configuration
|
||||
|
||||
app_mode = application_generate_entity.app_config.app_mode
|
||||
credit_deduction_metadata: dict[str, object] = {
|
||||
"provider": provider_name,
|
||||
"model": model_config.model,
|
||||
"model_type": ModelType.LLM.value,
|
||||
"app_type": get_credit_usage_app_type(app_mode),
|
||||
"created_by": get_credit_usage_created_by(app_mode),
|
||||
}
|
||||
credit_deduction_context: _CreditDeductionContext = {
|
||||
"request_id": str(message.id) if message.id else None,
|
||||
"metadata": credit_deduction_metadata,
|
||||
}
|
||||
agent_gateway_metered = (
|
||||
isinstance(application_generate_entity, AgentAppGenerateEntity)
|
||||
and application_generate_entity.agent_llm_gateway_enabled
|
||||
@@ -150,12 +170,14 @@ def handle(sender: Message, **kwargs):
|
||||
tenant_id=tenant_id,
|
||||
credits_required=used_quota,
|
||||
pool_type="trial",
|
||||
**credit_deduction_context,
|
||||
)
|
||||
case ProviderQuotaType.PAID:
|
||||
_deduct_credit_pool_quota_capped(
|
||||
tenant_id=tenant_id,
|
||||
credits_required=used_quota,
|
||||
pool_type="paid",
|
||||
**credit_deduction_context,
|
||||
)
|
||||
case ProviderQuotaType.FREE:
|
||||
quota_update = _ProviderUpdateOperation(
|
||||
@@ -205,16 +227,33 @@ def handle(sender: Message, **kwargs):
|
||||
raise
|
||||
|
||||
|
||||
def _deduct_credit_pool_quota_capped(*, tenant_id: str, credits_required: int, pool_type: str) -> None:
|
||||
def _deduct_credit_pool_quota_capped(
|
||||
*,
|
||||
tenant_id: str,
|
||||
credits_required: int,
|
||||
pool_type: str,
|
||||
request_id: str | None = None,
|
||||
metadata: Mapping[str, object] | None = None,
|
||||
) -> None:
|
||||
"""Apply post-generation credit accounting without failing message persistence on quota exhaustion."""
|
||||
from services.credit_pool_service import CreditPoolService
|
||||
|
||||
deducted_credits = CreditPoolService.deduct_credits_capped(
|
||||
tenant_id=tenant_id,
|
||||
credits_required=credits_required,
|
||||
pool_type=pool_type,
|
||||
session=db.session(),
|
||||
)
|
||||
if request_id is None and metadata is None:
|
||||
deducted_credits = CreditPoolService.deduct_credits_capped(
|
||||
tenant_id=tenant_id,
|
||||
credits_required=credits_required,
|
||||
pool_type=pool_type,
|
||||
session=db.session(),
|
||||
)
|
||||
else:
|
||||
deducted_credits = CreditPoolService.deduct_credits_capped(
|
||||
tenant_id=tenant_id,
|
||||
credits_required=credits_required,
|
||||
pool_type=pool_type,
|
||||
request_id=request_id,
|
||||
metadata=metadata,
|
||||
session=db.session(),
|
||||
)
|
||||
if deducted_credits < credits_required:
|
||||
logger.warning(
|
||||
"Credit pool exhausted during message-created accounting, "
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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"]
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
):
|
||||
|
||||
@@ -16,6 +16,7 @@ from collections.abc import Iterator
|
||||
from typing import Any
|
||||
|
||||
from core.app.app_config.entities import ModelConfig
|
||||
from core.credit_usage import CreditUsageAppType, CreditUsageCreatedBy
|
||||
from core.model_manager import ModelInstance, ModelManager
|
||||
from core.workflow.generator import WorkflowGenerator
|
||||
from core.workflow.generator.tool_catalogue import (
|
||||
@@ -67,7 +68,7 @@ class WorkflowGeneratorService:
|
||||
envelope as ``/rule-generate``).
|
||||
"""
|
||||
model_instance, model_parameters, tool_catalogue_entries, tool_catalogue_text, installed_tools = (
|
||||
cls._resolve_generation_context(tenant_id=tenant_id, model_config=model_config)
|
||||
cls._resolve_generation_context(tenant_id=tenant_id, mode=mode, model_config=model_config)
|
||||
)
|
||||
|
||||
return WorkflowGenerator.generate_workflow_graph(
|
||||
@@ -107,7 +108,7 @@ class WorkflowGeneratorService:
|
||||
single ``result`` SSE event).
|
||||
"""
|
||||
model_instance, model_parameters, tool_catalogue_entries, tool_catalogue_text, installed_tools = (
|
||||
cls._resolve_generation_context(tenant_id=tenant_id, model_config=model_config)
|
||||
cls._resolve_generation_context(tenant_id=tenant_id, mode=mode, model_config=model_config)
|
||||
)
|
||||
|
||||
yield from WorkflowGenerator.generate_workflow_graph_stream(
|
||||
@@ -130,6 +131,7 @@ class WorkflowGeneratorService:
|
||||
cls,
|
||||
*,
|
||||
tenant_id: str,
|
||||
mode: WorkflowGenerationModeRequest,
|
||||
model_config: ModelConfig,
|
||||
) -> tuple[
|
||||
ModelInstance,
|
||||
@@ -149,7 +151,17 @@ class WorkflowGeneratorService:
|
||||
empty set, so we don't reject every tool node just because we couldn't
|
||||
enumerate the catalogue).
|
||||
"""
|
||||
model_manager = ModelManager.for_tenant(tenant_id=tenant_id)
|
||||
app_type = {
|
||||
"workflow": CreditUsageAppType.WORKFLOW,
|
||||
"advanced-chat": CreditUsageAppType.CHATFLOW,
|
||||
}.get(mode, CreditUsageAppType.UNKNOWN)
|
||||
model_manager = ModelManager.for_tenant(
|
||||
tenant_id=tenant_id,
|
||||
request_metadata={
|
||||
"app_type": app_type,
|
||||
"created_by": CreditUsageCreatedBy.WORKFLOW_GENERATION,
|
||||
},
|
||||
)
|
||||
model_instance = model_manager.get_model_instance(
|
||||
tenant_id=tenant_id,
|
||||
model_type=ModelType.LLM,
|
||||
|
||||
@@ -10,7 +10,9 @@ import pandas as pd
|
||||
from celery import shared_task
|
||||
from sqlalchemy import func, select
|
||||
|
||||
from core.credit_usage import CreditUsageCreatedBy
|
||||
from core.db.session_factory import session_factory
|
||||
from core.model_context import with_credit_usage_created_by
|
||||
from core.model_manager import ModelManager
|
||||
from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType
|
||||
from extensions.ext_redis import redis_client
|
||||
@@ -27,6 +29,7 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@shared_task(queue="dataset")
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_INDEXING)
|
||||
def batch_create_segment_to_index_task(
|
||||
job_id: str,
|
||||
upload_file_id: str,
|
||||
|
||||
@@ -8,7 +8,9 @@ import click
|
||||
from celery import shared_task
|
||||
from sqlalchemy import or_, select
|
||||
|
||||
from core.credit_usage import CreditUsageCreatedBy
|
||||
from core.db.session_factory import session_factory
|
||||
from core.model_context import with_credit_usage_created_by
|
||||
from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType
|
||||
from models.dataset import Dataset, DocumentSegment, DocumentSegmentSummary
|
||||
from models.dataset import Document as DatasetDocument
|
||||
@@ -19,6 +21,7 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@shared_task(queue="dataset_summary")
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_INDEXING)
|
||||
def regenerate_summary_index_task(
|
||||
dataset_id: str,
|
||||
regenerate_reason: str = "summary_model_changed",
|
||||
|
||||
@@ -12,6 +12,8 @@ from sqlalchemy.orm import Session
|
||||
from core.agent.cot_agent_runner import CotAgentRunner
|
||||
from core.agent.entities import AgentScratchpadUnit
|
||||
from core.agent.errors import AgentMaxIterationError
|
||||
from core.app.entities.app_invoke_entities import CreditUsageCreatedBy
|
||||
from core.credit_usage import CreditUsageAppType
|
||||
from graphon.model_runtime.entities.llm_entities import LLMUsage
|
||||
from libs.datetime_utils import naive_utc_now
|
||||
from models.enums import ConversationFromSource, MessageStatus
|
||||
@@ -404,7 +406,11 @@ class TestRun:
|
||||
|
||||
results = list(runner.run(session, message, "query", {}))
|
||||
assert events == ["commit", "close", "first-chunk"]
|
||||
assert runner.model_instance.invoke_llm.call_args.kwargs["request_metadata"] == {"app_id": "app"}
|
||||
assert runner.model_instance.invoke_llm.call_args.kwargs["request_metadata"] == {
|
||||
"app_id": "app",
|
||||
"app_type": CreditUsageAppType.AGENT,
|
||||
"created_by": CreditUsageCreatedBy.APP.value,
|
||||
}
|
||||
assert results[-1].delta.message.content == ""
|
||||
|
||||
def test_run_usage_missing_key_branch(self, runner: DummyRunner, mocker: MockerFixture):
|
||||
|
||||
@@ -14,7 +14,9 @@ from sqlalchemy.orm import Session
|
||||
from core.agent.errors import AgentMaxIterationError
|
||||
from core.agent.fc_agent_runner import FunctionCallAgentRunner
|
||||
from core.app.apps.base_app_queue_manager import PublishFrom
|
||||
from core.app.entities.app_invoke_entities import CreditUsageCreatedBy
|
||||
from core.app.entities.queue_entities import QueueMessageFileEvent
|
||||
from core.credit_usage import CreditUsageAppType
|
||||
from graphon.model_runtime.entities.llm_entities import LLMUsage
|
||||
from graphon.model_runtime.entities.message_entities import (
|
||||
DocumentPromptMessageContent,
|
||||
@@ -435,7 +437,11 @@ class TestRunMethod:
|
||||
assert len(outputs) == 1
|
||||
assert "session" not in runner.create_agent_thought.call_args.kwargs
|
||||
assert "session" not in runner.save_agent_thought.call_args.kwargs
|
||||
assert runner.model_instance.invoke_llm.call_args.kwargs["request_metadata"] == {"app_id": "app"}
|
||||
assert runner.model_instance.invoke_llm.call_args.kwargs["request_metadata"] == {
|
||||
"app_id": "app",
|
||||
"app_type": CreditUsageAppType.AGENT,
|
||||
"created_by": CreditUsageCreatedBy.APP.value,
|
||||
}
|
||||
runner.queue_manager.publish.assert_called()
|
||||
|
||||
queue_calls = runner.queue_manager.publish.call_args_list
|
||||
|
||||
@@ -269,6 +269,7 @@ class TestChatAppRunner:
|
||||
app_config = SimpleNamespace(
|
||||
app_id="app-1",
|
||||
tenant_id="tenant-1",
|
||||
app_mode=AppMode.CHAT,
|
||||
prompt_template=None,
|
||||
external_data_variables=[],
|
||||
dataset=None,
|
||||
@@ -311,6 +312,7 @@ class TestChatAppRunner:
|
||||
app_config = SimpleNamespace(
|
||||
app_id="app-1",
|
||||
tenant_id="tenant-1",
|
||||
app_mode=AppMode.CHAT,
|
||||
prompt_template=None,
|
||||
external_data_variables=[],
|
||||
dataset=None,
|
||||
@@ -394,6 +396,7 @@ class TestChatAppRunner:
|
||||
app_config = SimpleNamespace(
|
||||
app_id="app-1",
|
||||
tenant_id="tenant-1",
|
||||
app_mode=AppMode.CHAT,
|
||||
prompt_template=None,
|
||||
external_data_variables=[],
|
||||
dataset=None,
|
||||
|
||||
@@ -7,6 +7,7 @@ from sqlalchemy.orm import Session
|
||||
|
||||
import core.app.apps.completion.app_runner as module
|
||||
from core.app.apps.completion.app_runner import CompletionAppRunner
|
||||
from core.credit_usage import CreditUsageAppType, CreditUsageCreatedBy
|
||||
from core.moderation.base import ModerationError
|
||||
from graphon.model_runtime.entities.message_entities import ImagePromptMessageContent
|
||||
from graphon.model_runtime.entities.model_entities import ModelType
|
||||
@@ -25,6 +26,7 @@ def _build_app_config(dataset=None, external_tools=None, additional_features=Non
|
||||
app_config = MagicMock()
|
||||
app_config.app_id = APP_ID
|
||||
app_config.tenant_id = TENANT_ID
|
||||
app_config.app_mode = AppMode.COMPLETION
|
||||
app_config.prompt_template = MagicMock()
|
||||
app_config.dataset = dataset
|
||||
app_config.external_data_variables = external_tools or []
|
||||
@@ -245,7 +247,11 @@ class TestCompletionAppRunner:
|
||||
model_parameters={"max_tokens": 10},
|
||||
stop=["stop"],
|
||||
stream=stream,
|
||||
request_metadata={"app_id": APP_ID},
|
||||
request_metadata={
|
||||
"app_id": APP_ID,
|
||||
"app_type": CreditUsageAppType.COMPLETION,
|
||||
"created_by": CreditUsageCreatedBy.APP,
|
||||
},
|
||||
)
|
||||
|
||||
def test_run_uses_low_image_detail_default(self, runner, mocker: MockerFixture, sqlite_session: Session):
|
||||
|
||||
@@ -10,9 +10,11 @@ from core.app.apps.base_app_queue_manager import AppQueueManager
|
||||
from core.app.apps.workflow.app_runner import WorkflowAppRunner
|
||||
from core.app.apps.workflow_app_runner import WorkflowBasedAppRunner
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom, WorkflowAppGenerateEntity
|
||||
from core.credit_usage import CreditUsageAppType
|
||||
from core.workflow.system_variables import default_system_variables
|
||||
from graphon.entities.graph_config import NodeConfigDictAdapter
|
||||
from graphon.runtime import GraphRuntimeState, VariablePool
|
||||
from models.model import AppMode
|
||||
from models.workflow import Workflow, WorkflowKind
|
||||
|
||||
|
||||
@@ -41,6 +43,7 @@ def test_run_uses_single_node_execution_branch(
|
||||
app_config.app_id = "app"
|
||||
app_config.tenant_id = "tenant"
|
||||
app_config.workflow_id = "workflow"
|
||||
app_config.app_mode = AppMode.WORKFLOW
|
||||
|
||||
app_generate_entity = MagicMock(spec=WorkflowAppGenerateEntity)
|
||||
app_generate_entity.app_config = app_config
|
||||
@@ -104,6 +107,7 @@ def test_run_uses_single_node_execution_branch(
|
||||
single_iteration_run=single_iteration_run,
|
||||
single_loop_run=single_loop_run,
|
||||
user_id="user",
|
||||
app_type=CreditUsageAppType.WORKFLOW,
|
||||
trace_session_id="session-1",
|
||||
)
|
||||
init_graph.assert_not_called()
|
||||
|
||||
@@ -7,14 +7,19 @@ import pytest
|
||||
from core.app.app_config.entities import WorkflowUIBasedAppConfig
|
||||
from core.app.entities.app_invoke_entities import (
|
||||
AdvancedChatAppGenerateEntity,
|
||||
CreditUsageCreatedBy,
|
||||
DifyRunContext,
|
||||
InvokeFrom,
|
||||
WorkflowAppGenerateEntity,
|
||||
get_credit_usage_app_type,
|
||||
get_credit_usage_created_by,
|
||||
)
|
||||
from core.app.layers.pause_state_persist_layer import (
|
||||
WorkflowResumptionContext,
|
||||
_AdvancedChatAppGenerateEntityWrapper,
|
||||
_WorkflowGenerateEntityWrapper,
|
||||
)
|
||||
from core.credit_usage import CreditUsageAppType
|
||||
from core.ops.ops_trace_manager import TraceQueueManager
|
||||
from models.model import AppMode
|
||||
|
||||
@@ -99,6 +104,60 @@ def test_advanced_chat_generate_entity_roundtrip_excludes_trace_manager():
|
||||
assert restored.trace_manager is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("app_mode", "created_by"),
|
||||
[
|
||||
(AppMode.CHAT, CreditUsageCreatedBy.APP),
|
||||
(AppMode.ADVANCED_CHAT, CreditUsageCreatedBy.APP),
|
||||
(AppMode.WORKFLOW, CreditUsageCreatedBy.APP),
|
||||
(AppMode.AGENT_CHAT, CreditUsageCreatedBy.APP),
|
||||
(AppMode.AGENT, CreditUsageCreatedBy.APP),
|
||||
(AppMode.COMPLETION, CreditUsageCreatedBy.APP),
|
||||
(AppMode.RAG_PIPELINE, CreditUsageCreatedBy.APP),
|
||||
],
|
||||
)
|
||||
def test_get_credit_usage_created_by_maps_app_mode(
|
||||
app_mode: AppMode,
|
||||
created_by: CreditUsageCreatedBy,
|
||||
) -> None:
|
||||
assert get_credit_usage_created_by(app_mode) == created_by
|
||||
|
||||
|
||||
def test_get_credit_usage_created_by_defaults_to_unknown() -> None:
|
||||
assert get_credit_usage_created_by(None) == CreditUsageCreatedBy.UNKNOWN
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("app_mode", "app_type"),
|
||||
[
|
||||
(AppMode.CHAT, CreditUsageAppType.CHATBOT),
|
||||
(AppMode.ADVANCED_CHAT, CreditUsageAppType.CHATFLOW),
|
||||
(AppMode.WORKFLOW, CreditUsageAppType.WORKFLOW),
|
||||
(AppMode.AGENT_CHAT, CreditUsageAppType.AGENT),
|
||||
(AppMode.AGENT, CreditUsageAppType.AGENT_V2),
|
||||
(AppMode.COMPLETION, CreditUsageAppType.COMPLETION),
|
||||
(AppMode.RAG_PIPELINE, CreditUsageAppType.RAG_PIPELINE),
|
||||
],
|
||||
)
|
||||
def test_get_credit_usage_app_type_maps_app_mode(app_mode: AppMode, app_type: CreditUsageAppType) -> None:
|
||||
assert get_credit_usage_app_type(app_mode) == app_type
|
||||
|
||||
|
||||
def test_dify_run_context_normalizes_unknown_created_by_values() -> None:
|
||||
context = DifyRunContext(
|
||||
tenant_id="tenant-id",
|
||||
app_id="app-id",
|
||||
user_id="user-id",
|
||||
user_from="account",
|
||||
invoke_from=InvokeFrom.DEBUGGER,
|
||||
app_type="legacy-app-type",
|
||||
created_by="legacy-created-by",
|
||||
)
|
||||
|
||||
assert context.app_type == CreditUsageAppType.UNKNOWN
|
||||
assert context.created_by == CreditUsageCreatedBy.UNKNOWN
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ResumptionContextCase:
|
||||
name: str
|
||||
|
||||
+4
-4
@@ -1043,7 +1043,7 @@ class TestEasyUiBasedGenerateTaskPipeline:
|
||||
_set_method(pipeline, "_process_stream_response", lambda publisher, trace_manager: iter([payload]))
|
||||
monkeypatch.setattr(
|
||||
"core.app.task_pipeline.easy_ui_based_generate_task_pipeline.AppGeneratorTTSPublisher",
|
||||
lambda tenant_id, voice, language: _Publisher(),
|
||||
lambda tenant_id, voice, language, app_type: _Publisher(),
|
||||
)
|
||||
|
||||
responses = list(pipeline._wrapper_process_stream_response())
|
||||
@@ -1085,7 +1085,7 @@ class TestEasyUiBasedGenerateTaskPipeline:
|
||||
_set_method(pipeline, "_process_stream_response", lambda publisher, trace_manager: iter([]))
|
||||
monkeypatch.setattr(
|
||||
"core.app.task_pipeline.easy_ui_based_generate_task_pipeline.AppGeneratorTTSPublisher",
|
||||
lambda tenant_id, voice, language: _Publisher(),
|
||||
lambda tenant_id, voice, language, app_type: _Publisher(),
|
||||
)
|
||||
responses = list(pipeline._wrapper_process_stream_response())
|
||||
|
||||
@@ -1124,7 +1124,7 @@ class TestEasyUiBasedGenerateTaskPipeline:
|
||||
_set_method(pipeline, "_process_stream_response", lambda publisher, trace_manager: iter([]))
|
||||
monkeypatch.setattr(
|
||||
"core.app.task_pipeline.easy_ui_based_generate_task_pipeline.AppGeneratorTTSPublisher",
|
||||
lambda tenant_id, voice, language: _Publisher(),
|
||||
lambda tenant_id, voice, language, app_type: _Publisher(),
|
||||
)
|
||||
|
||||
responses = list(pipeline._wrapper_process_stream_response())
|
||||
@@ -1153,7 +1153,7 @@ class TestEasyUiBasedGenerateTaskPipeline:
|
||||
_set_method(pipeline, "_process_stream_response", lambda publisher, trace_manager: iter([error]))
|
||||
monkeypatch.setattr(
|
||||
"core.app.task_pipeline.easy_ui_based_generate_task_pipeline.AppGeneratorTTSPublisher",
|
||||
lambda tenant_id, voice, language: publisher,
|
||||
lambda tenant_id, voice, language, app_type: publisher,
|
||||
)
|
||||
|
||||
responses = list(pipeline._wrapper_process_stream_response())
|
||||
|
||||
@@ -10,6 +10,7 @@ from sqlalchemy.engine import Engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from configs import dify_config
|
||||
from core.app.entities.app_invoke_entities import CreditUsageCreatedBy
|
||||
from core.app.llm.quota import (
|
||||
LLMQuotaReservationState,
|
||||
deduct_llm_quota,
|
||||
@@ -19,6 +20,7 @@ from core.app.llm.quota import (
|
||||
reserve_llm_quota_for_model,
|
||||
reserve_model_quota_for_model,
|
||||
)
|
||||
from core.credit_usage import CreditUsageAppType
|
||||
from core.entities.model_entities import ModelStatus
|
||||
from core.entities.provider_entities import ProviderQuotaType, QuotaUnit
|
||||
from core.errors.error import QuotaExceededError
|
||||
@@ -133,6 +135,8 @@ def test_reserve_llm_quota_uses_exact_credit_pool_reservation() -> None:
|
||||
provider="openai",
|
||||
model="gpt-4o",
|
||||
request_id="11111111-1111-5111-8111-111111111111",
|
||||
app_type=CreditUsageAppType.CHATBOT,
|
||||
created_by=CreditUsageCreatedBy.APP.value,
|
||||
)
|
||||
reservation.commit(LLMUsage.empty_usage())
|
||||
reservation.release()
|
||||
@@ -145,7 +149,13 @@ def test_reserve_llm_quota_uses_exact_credit_pool_reservation() -> None:
|
||||
pool_type="trial",
|
||||
request_id="11111111-1111-5111-8111-111111111111",
|
||||
session_factory=ANY,
|
||||
meta={"source": "llm.invoke", "provider": "openai", "model": "gpt-4o"},
|
||||
meta={
|
||||
"source": "llm.invoke",
|
||||
"provider": "openai",
|
||||
"model": "gpt-4o",
|
||||
"app_type": CreditUsageAppType.CHATBOT,
|
||||
"created_by": CreditUsageCreatedBy.APP.value,
|
||||
},
|
||||
)
|
||||
credit_reservation.commit.assert_called_once_with()
|
||||
credit_reservation.release.assert_not_called()
|
||||
@@ -227,6 +237,8 @@ def test_reserve_non_llm_quota_uses_model_type_and_credit_pool_reservation() ->
|
||||
"provider": "openai",
|
||||
"model_type": "text-embedding",
|
||||
"model": "text-embedding-3-small",
|
||||
"app_type": CreditUsageAppType.UNKNOWN,
|
||||
"created_by": "unknown",
|
||||
},
|
||||
)
|
||||
credit_reservation.commit.assert_called_once_with()
|
||||
@@ -354,6 +366,13 @@ def test_deduct_llm_quota_for_model_uses_identity_based_trial_billing() -> None:
|
||||
mock_deduct_credits.assert_called_once_with(
|
||||
tenant_id="tenant-id",
|
||||
credits_required=42,
|
||||
metadata={
|
||||
"provider": "openai",
|
||||
"model": "gpt-4o",
|
||||
"model_type": "llm",
|
||||
"app_type": CreditUsageAppType.UNKNOWN,
|
||||
"created_by": "unknown",
|
||||
},
|
||||
session=ANY,
|
||||
)
|
||||
|
||||
@@ -474,6 +493,13 @@ def test_deduct_llm_quota_for_model_uses_credit_configuration() -> None:
|
||||
mock_deduct_credits.assert_called_once_with(
|
||||
tenant_id="tenant-id",
|
||||
credits_required=9,
|
||||
metadata={
|
||||
"provider": "openai",
|
||||
"model": "gpt-4o",
|
||||
"model_type": "llm",
|
||||
"app_type": CreditUsageAppType.UNKNOWN,
|
||||
"created_by": "unknown",
|
||||
},
|
||||
session=ANY,
|
||||
)
|
||||
|
||||
@@ -510,6 +536,13 @@ def test_deduct_llm_quota_for_model_uses_single_charge_for_times_quota() -> None
|
||||
mock_deduct_credits.assert_called_once_with(
|
||||
tenant_id="tenant-id",
|
||||
credits_required=1,
|
||||
metadata={
|
||||
"provider": "openai",
|
||||
"model": "gpt-4o",
|
||||
"model_type": "llm",
|
||||
"app_type": CreditUsageAppType.UNKNOWN,
|
||||
"created_by": "unknown",
|
||||
},
|
||||
session=ANY,
|
||||
)
|
||||
|
||||
@@ -548,6 +581,13 @@ def test_deduct_llm_quota_for_model_uses_paid_billing_pool() -> None:
|
||||
tenant_id="tenant-id",
|
||||
credits_required=5,
|
||||
pool_type="paid",
|
||||
metadata={
|
||||
"provider": "openai",
|
||||
"model": "gpt-4o",
|
||||
"model_type": "llm",
|
||||
"app_type": CreditUsageAppType.UNKNOWN,
|
||||
"created_by": "unknown",
|
||||
},
|
||||
session=ANY,
|
||||
)
|
||||
|
||||
|
||||
@@ -7,12 +7,14 @@ from core.moderation.base import ModerationAction, ModerationError, ModerationIn
|
||||
from core.moderation.input_moderation import InputModeration
|
||||
from core.ops.entities.trace_entity import TraceTaskName
|
||||
from core.ops.ops_trace_manager import TraceQueueManager
|
||||
from models.model import AppMode
|
||||
|
||||
|
||||
class TestInputModeration:
|
||||
@pytest.fixture
|
||||
def app_config(self):
|
||||
config = MagicMock(spec=AppConfig)
|
||||
config.app_mode = AppMode.CHAT
|
||||
config.sensitive_word_avoidance = None
|
||||
return config
|
||||
|
||||
|
||||
@@ -755,7 +755,7 @@ class TestIndexingRunnerLoad:
|
||||
# Verify executor was used for parallel processing
|
||||
assert mock_executor_instance.submit.called
|
||||
for submit_call in mock_executor_instance.submit.call_args_list:
|
||||
assert submit_call.args[0] == runner._process_chunk
|
||||
assert submit_call.args[0].__name__ == runner._process_chunk.__name__
|
||||
assert len(submit_call.args) == 6
|
||||
mock_future.result.assert_called()
|
||||
assert mock_update_status.call_args.kwargs["extra_update_params"][DatasetDocument.tokens] == 300
|
||||
|
||||
@@ -0,0 +1,33 @@
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
from core.credit_usage import CreditUsageAppType, CreditUsageCreatedBy
|
||||
from core.model_context import get_credit_usage_metadata, use_credit_usage_metadata, with_credit_usage_created_by
|
||||
from extensions.otel import propagate_context
|
||||
|
||||
|
||||
def test_credit_usage_metadata_is_scoped_and_propagated_to_worker_threads() -> None:
|
||||
with use_credit_usage_metadata({"app_type": CreditUsageAppType.WORKFLOW}):
|
||||
assert get_credit_usage_metadata() == {"app_type": CreditUsageAppType.WORKFLOW}
|
||||
|
||||
with ThreadPoolExecutor(max_workers=1) as executor:
|
||||
read_context = propagate_context(get_credit_usage_metadata)
|
||||
assert executor.submit(read_context).result() == {"app_type": CreditUsageAppType.WORKFLOW}
|
||||
|
||||
assert get_credit_usage_metadata() is None
|
||||
|
||||
|
||||
def test_nested_credit_usage_metadata_keeps_app_type_and_direct_feature() -> None:
|
||||
with use_credit_usage_metadata({"app_type": CreditUsageAppType.WORKFLOW}):
|
||||
with use_credit_usage_metadata({"created_by": CreditUsageCreatedBy.KNOWLEDGE_RETRIEVAL}):
|
||||
assert get_credit_usage_metadata() == {
|
||||
"app_type": CreditUsageAppType.WORKFLOW,
|
||||
"created_by": CreditUsageCreatedBy.KNOWLEDGE_RETRIEVAL,
|
||||
}
|
||||
|
||||
|
||||
def test_credit_usage_created_by_decorator_sets_feature() -> None:
|
||||
@with_credit_usage_created_by(CreditUsageCreatedBy.KNOWLEDGE_RETRIEVAL)
|
||||
def read_context() -> object:
|
||||
return get_credit_usage_metadata()
|
||||
|
||||
assert read_context() == {"created_by": CreditUsageCreatedBy.KNOWLEDGE_RETRIEVAL}
|
||||
@@ -7,6 +7,7 @@ import pytest
|
||||
import redis
|
||||
from pytest_mock import MockerFixture
|
||||
|
||||
from core.app.entities.app_invoke_entities import CreditUsageCreatedBy
|
||||
from core.entities.provider_entities import (
|
||||
ModelLoadBalancingConfiguration,
|
||||
ProviderQuotaType,
|
||||
@@ -184,11 +185,14 @@ def test_quota_managed_non_streaming_invocation_finalizes_reservation() -> None:
|
||||
response = model_instance.invoke_llm(
|
||||
prompt_messages=[],
|
||||
stream=False,
|
||||
request_metadata={"invocation_id": invocation_id},
|
||||
request_metadata={"invocation_id": invocation_id, "created_by": CreditUsageCreatedBy.APP.value},
|
||||
)
|
||||
|
||||
assert response is result
|
||||
reserve_quota.assert_called_once_with(request_id=invocation_id)
|
||||
reserve_quota.assert_called_once_with(
|
||||
request_id=invocation_id,
|
||||
created_by=CreditUsageCreatedBy.APP.value,
|
||||
)
|
||||
invoke.assert_called_once()
|
||||
reservation.commit.assert_called_once_with(usage)
|
||||
reservation.release.assert_called_once_with()
|
||||
|
||||
@@ -481,6 +481,8 @@ class TestDifyNodeFactoryCreateNode:
|
||||
app_id="app-id",
|
||||
user_id="user-id",
|
||||
invoke_from=InvokeFrom.DEBUGGER,
|
||||
app_type=None,
|
||||
created_by=None,
|
||||
)
|
||||
factory._code_executor = sentinel.code_executor
|
||||
factory._code_limits = sentinel.code_limits
|
||||
|
||||
@@ -277,7 +277,7 @@ def test_dify_prepared_llm_requires_model_schema() -> None:
|
||||
|
||||
def test_dify_prepared_llm_delegates_structured_output_helper(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
model_instance = _ModelInstanceStub(model_schema=_build_model_schema())
|
||||
prepared = DifyPreparedLLM(model_instance)
|
||||
prepared = DifyPreparedLLM(model_instance, request_metadata={"created_by": "workflow"})
|
||||
invoke_structured = MagicMock(return_value=sentinel.structured)
|
||||
monkeypatch.setattr(node_runtime, "invoke_llm_with_structured_output", invoke_structured)
|
||||
|
||||
@@ -299,6 +299,7 @@ def test_dify_prepared_llm_delegates_structured_output_helper(monkeypatch: pytes
|
||||
model_parameters={"temperature": 0.2},
|
||||
stop=["done"],
|
||||
stream=True,
|
||||
request_metadata={"created_by": "workflow"},
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -8,6 +8,7 @@ import pytest
|
||||
|
||||
from core.app.apps.exc import GenerateTaskStoppedError
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom, UserFrom
|
||||
from core.credit_usage import CreditUsageAppType
|
||||
from core.workflow import workflow_entry
|
||||
from core.workflow.system_variables import default_system_variables
|
||||
from graphon.entities.base_node_data import BaseNodeData
|
||||
@@ -610,6 +611,7 @@ class TestWorkflowEntryHelpers:
|
||||
user_id="user-id",
|
||||
user_from=UserFrom.ACCOUNT,
|
||||
invoke_from=InvokeFrom.DEBUGGER,
|
||||
app_type=CreditUsageAppType.WORKFLOW,
|
||||
)
|
||||
graph_init_context_cls.assert_called_once_with(
|
||||
workflow_id="",
|
||||
|
||||
@@ -7,10 +7,11 @@ import pytest
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
|
||||
from core.app.entities.app_invoke_entities import AgentAppGenerateEntity, ChatAppGenerateEntity
|
||||
from core.app.entities.app_invoke_entities import AgentAppGenerateEntity, ChatAppGenerateEntity, CreditUsageCreatedBy
|
||||
from core.credit_usage import CreditUsageAppType
|
||||
from core.entities.provider_entities import ProviderQuotaType, QuotaUnit
|
||||
from events.event_handlers import update_provider_when_message_created
|
||||
from models import Message, TenantCreditPool
|
||||
from models import AppMode, Message, TenantCreditPool
|
||||
from models.enums import ProviderQuotaType as ModelProviderQuotaType
|
||||
from models.provider import ProviderType
|
||||
|
||||
@@ -48,7 +49,7 @@ def test_message_created_trial_credit_accounting_does_not_raise_when_balance_is_
|
||||
],
|
||||
)
|
||||
application_generate_entity = ChatAppGenerateEntity.model_construct(
|
||||
app_config=SimpleNamespace(tenant_id=tenant_id),
|
||||
app_config=SimpleNamespace(tenant_id=tenant_id, app_mode=AppMode.CHAT),
|
||||
model_conf=SimpleNamespace(
|
||||
provider="openai",
|
||||
model="gpt-4o",
|
||||
@@ -61,6 +62,7 @@ def test_message_created_trial_credit_accounting_does_not_raise_when_balance_is_
|
||||
),
|
||||
)
|
||||
message = Message(message_tokens=2, answer_tokens=1)
|
||||
message.id = "message-1"
|
||||
|
||||
with (
|
||||
patch.object(update_provider_when_message_created, "_execute_provider_updates"),
|
||||
@@ -89,7 +91,7 @@ def test_message_created_paid_credit_accounting_uses_paid_pool() -> None:
|
||||
],
|
||||
)
|
||||
application_generate_entity = ChatAppGenerateEntity.model_construct(
|
||||
app_config=SimpleNamespace(tenant_id=tenant_id),
|
||||
app_config=SimpleNamespace(tenant_id=tenant_id, app_mode=AppMode.CHAT),
|
||||
model_conf=SimpleNamespace(
|
||||
provider="openai",
|
||||
model="gpt-4o",
|
||||
@@ -102,6 +104,7 @@ def test_message_created_paid_credit_accounting_uses_paid_pool() -> None:
|
||||
),
|
||||
)
|
||||
message = Message(message_tokens=2, answer_tokens=1)
|
||||
message.id = "message-1"
|
||||
|
||||
with (
|
||||
patch.object(update_provider_when_message_created, "_deduct_credit_pool_quota_capped") as mock_deduct,
|
||||
@@ -116,6 +119,14 @@ def test_message_created_paid_credit_accounting_uses_paid_pool() -> None:
|
||||
tenant_id=tenant_id,
|
||||
credits_required=3,
|
||||
pool_type="paid",
|
||||
request_id="message-1",
|
||||
metadata={
|
||||
"provider": "openai",
|
||||
"model": "gpt-4o",
|
||||
"model_type": "llm",
|
||||
"app_type": CreditUsageAppType.CHATBOT,
|
||||
"created_by": CreditUsageCreatedBy.APP.value,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@@ -132,7 +143,7 @@ def test_agent_app_gateway_accounting_skips_legacy_message_charge() -> None:
|
||||
],
|
||||
)
|
||||
application_generate_entity = AgentAppGenerateEntity.model_construct(
|
||||
app_config=SimpleNamespace(tenant_id=tenant_id),
|
||||
app_config=SimpleNamespace(tenant_id=tenant_id, app_mode=AppMode.AGENT),
|
||||
model_conf=SimpleNamespace(
|
||||
provider="openai",
|
||||
model="gpt-4o",
|
||||
|
||||
@@ -11,6 +11,7 @@ from pydantic import JsonValue
|
||||
from sqlalchemy.orm import Session, sessionmaker
|
||||
|
||||
from configs import dify_config
|
||||
from core.credit_usage import CreditUsageAppType, CreditUsageCreatedBy
|
||||
from core.entities.model_entities import ModelStatus
|
||||
from core.entities.provider_entities import ProviderQuotaType, QuotaUnit
|
||||
from core.model_manager import ModelInstance, QuotaManagedModelInstance
|
||||
@@ -63,6 +64,7 @@ def _model_instance(
|
||||
instance.credentials = {"api_key": "hosted"}
|
||||
instance.model_type_instance = MagicMock()
|
||||
instance.load_balancing_manager = None
|
||||
instance._request_metadata = None
|
||||
return instance, provider_model
|
||||
|
||||
|
||||
@@ -71,13 +73,14 @@ def _persist_app(
|
||||
*,
|
||||
request: AgentLLMInvokeRequest,
|
||||
tenant_id: str | None = None,
|
||||
mode: AppMode = AppMode.CHAT,
|
||||
) -> App:
|
||||
app = App(
|
||||
id=request.caller.app_id,
|
||||
tenant_id=tenant_id or request.caller.tenant_id,
|
||||
name="Agent LLM gateway test app",
|
||||
description="",
|
||||
mode=AppMode.CHAT,
|
||||
mode=mode,
|
||||
enable_site=False,
|
||||
enable_api=False,
|
||||
max_active_requests=None,
|
||||
@@ -175,7 +178,7 @@ def test_gateway_uses_quota_managed_instance_as_single_credit_owner(
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
request = _request()
|
||||
_persist_app(sqlite_session, request=request)
|
||||
_persist_app(sqlite_session, request=request, mode=AppMode.WORKFLOW)
|
||||
service = AgentLLMInnerService(session_factory=sqlite_session_factory)
|
||||
model_instance, _ = _model_instance()
|
||||
reservation = MagicMock(commit_before_delivery=True)
|
||||
@@ -190,7 +193,12 @@ def test_gateway_uses_quota_managed_instance_as_single_credit_owner(
|
||||
chunks = list(service.invoke(prepared))
|
||||
|
||||
assert chunks == [provider_chunk]
|
||||
model_instance.reserve_quota.assert_called_once_with(request_id=request.caller.invocation_id)
|
||||
model_instance.reserve_quota.assert_called_once_with(
|
||||
request_id=request.caller.invocation_id,
|
||||
app_type=CreditUsageAppType.WORKFLOW,
|
||||
created_by=CreditUsageCreatedBy.AGENT_NODE,
|
||||
)
|
||||
assert provider_invoke.call_args.kwargs["request_metadata"]["agent_config_version_kind"] == "draft"
|
||||
reservation.commit.assert_called_once_with(provider_chunk.delta.usage)
|
||||
reservation.release.assert_called_once_with()
|
||||
provider_invoke.assert_called_once()
|
||||
@@ -227,6 +235,39 @@ def test_gateway_forwards_prompt_messages_without_revalidation() -> None:
|
||||
assert model_instance.invoke_llm.call_args.kwargs["prompt_messages"] is prompt_messages
|
||||
|
||||
|
||||
def test_standalone_agent_app_gateway_is_attributed_to_app() -> None:
|
||||
request = _request()
|
||||
model_instance = MagicMock(spec=ModelInstance)
|
||||
model_instance.invoke_llm.return_value = iter([])
|
||||
prepared = PreparedAgentLLMInvocation(
|
||||
request=request,
|
||||
model_instance=model_instance,
|
||||
app_type=CreditUsageAppType.AGENT_V2,
|
||||
)
|
||||
|
||||
assert list(AgentLLMInnerService().invoke(prepared)) == []
|
||||
request_metadata = model_instance.invoke_llm.call_args.kwargs["request_metadata"]
|
||||
assert request_metadata["created_by"] is CreditUsageCreatedBy.APP
|
||||
assert request_metadata["agent_config_version_kind"] == "draft"
|
||||
|
||||
|
||||
def test_agent_build_draft_gateway_is_attributed_to_build_draft() -> None:
|
||||
request = _request()
|
||||
request.caller.agent_config_version_kind = "build_draft"
|
||||
model_instance = MagicMock(spec=ModelInstance)
|
||||
model_instance.invoke_llm.return_value = iter([])
|
||||
prepared = PreparedAgentLLMInvocation(
|
||||
request=request,
|
||||
model_instance=model_instance,
|
||||
app_type=CreditUsageAppType.AGENT_V2,
|
||||
)
|
||||
|
||||
assert list(AgentLLMInnerService().invoke(prepared)) == []
|
||||
request_metadata = model_instance.invoke_llm.call_args.kwargs["request_metadata"]
|
||||
assert request_metadata["created_by"] is CreditUsageCreatedBy.BUILD_DRAFT
|
||||
assert request_metadata["agent_config_version_kind"] == "build_draft"
|
||||
|
||||
|
||||
def test_retried_gateway_delivery_uses_one_effective_billing_charge(
|
||||
sqlite_session_factory: sessionmaker[Session],
|
||||
sqlite_session: Session,
|
||||
|
||||
@@ -64,6 +64,7 @@ from flask import Flask
|
||||
from sqlalchemy.orm import Session
|
||||
from werkzeug.datastructures import FileStorage
|
||||
|
||||
from core.credit_usage import CreditUsageAppType, CreditUsageCreatedBy
|
||||
from core.plugin.entities.plugin_daemon import TTSAudioChunk
|
||||
from graphon.model_runtime.errors.invoke import InvokeBadRequestError
|
||||
from models.agent_config_entities import AgentSoulConfig
|
||||
@@ -282,7 +283,14 @@ class TestAudioServiceASR:
|
||||
# Assert
|
||||
assert result == {"text": "Transcribed text"}
|
||||
mock_model_instance.invoke_speech2text.assert_called_once()
|
||||
mock_model_manager_class.assert_called_once_with(tenant_id=app.tenant_id, user_id="user-123")
|
||||
mock_model_manager_class.assert_called_once_with(
|
||||
tenant_id=app.tenant_id,
|
||||
user_id="user-123",
|
||||
request_metadata={
|
||||
"app_type": CreditUsageAppType.CHATBOT,
|
||||
"created_by": CreditUsageCreatedBy.AUDIO,
|
||||
},
|
||||
)
|
||||
|
||||
@patch("services.audio_service.ModelManager.for_tenant", autospec=True)
|
||||
def test_transcript_asr_accepts_x_m4a_mimetype(
|
||||
@@ -393,7 +401,14 @@ class TestAudioServiceASR:
|
||||
)
|
||||
|
||||
assert result == {"text": "Agent transcript"}
|
||||
mock_model_manager_class.assert_called_once_with(tenant_id=app.tenant_id, user_id="account-1")
|
||||
mock_model_manager_class.assert_called_once_with(
|
||||
tenant_id=app.tenant_id,
|
||||
user_id="account-1",
|
||||
request_metadata={
|
||||
"app_type": CreditUsageAppType.AGENT_V2,
|
||||
"created_by": CreditUsageCreatedBy.AUDIO,
|
||||
},
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"agent_soul",
|
||||
@@ -586,7 +601,14 @@ class TestAudioServiceTTS:
|
||||
# Assert
|
||||
assert result.content_type == "audio/mpeg"
|
||||
assert result.get_data() == b"audio data"
|
||||
mock_model_manager_class.assert_called_once_with(tenant_id=app.tenant_id, user_id="user-123")
|
||||
mock_model_manager_class.assert_called_once_with(
|
||||
tenant_id=app.tenant_id,
|
||||
user_id="user-123",
|
||||
request_metadata={
|
||||
"app_type": CreditUsageAppType.CHATBOT,
|
||||
"created_by": CreditUsageCreatedBy.AUDIO,
|
||||
},
|
||||
)
|
||||
mock_model_instance.invoke_tts.assert_called_once_with(
|
||||
content_text="Hello world",
|
||||
voice="en-US-Neural",
|
||||
|
||||
@@ -9,6 +9,8 @@ import pytest
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from core.app.entities.app_invoke_entities import CreditUsageCreatedBy
|
||||
from core.credit_usage import CreditUsageAppType
|
||||
from core.errors.error import QuotaExceededError
|
||||
from enums import DeploymentEdition
|
||||
from models import TenantCreditPool
|
||||
@@ -301,7 +303,11 @@ def test_reserve_credits_commits_billing_reservation_once() -> None:
|
||||
bucket="trial",
|
||||
request_id="request-1",
|
||||
amount=3,
|
||||
meta={"source": "test"},
|
||||
meta={
|
||||
"source": "test",
|
||||
"created_by": CreditUsageCreatedBy.UNKNOWN.value,
|
||||
"app_type": CreditUsageAppType.UNKNOWN.value,
|
||||
},
|
||||
)
|
||||
quota_commit.assert_called_once_with(
|
||||
tenant_id="tenant-1",
|
||||
@@ -309,7 +315,12 @@ def test_reserve_credits_commits_billing_reservation_once() -> None:
|
||||
bucket="trial",
|
||||
reservation_id="reservation-1",
|
||||
actual_amount=3,
|
||||
meta={"source": "test", "request_id": "request-1"},
|
||||
meta={
|
||||
"source": "test",
|
||||
"created_by": CreditUsageCreatedBy.UNKNOWN.value,
|
||||
"app_type": CreditUsageAppType.UNKNOWN.value,
|
||||
"request_id": "request-1",
|
||||
},
|
||||
)
|
||||
quota_release.assert_not_called()
|
||||
|
||||
@@ -381,7 +392,11 @@ def test_check_and_deduct_credits_uses_billing_reserve_and_commit_when_enabled()
|
||||
bucket="trial",
|
||||
request_id=ANY,
|
||||
amount=3,
|
||||
meta={"source": "credit_pool.check_and_deduct"},
|
||||
meta={
|
||||
"source": "credit_pool.check_and_deduct",
|
||||
"created_by": CreditUsageCreatedBy.UNKNOWN.value,
|
||||
"app_type": CreditUsageAppType.UNKNOWN.value,
|
||||
},
|
||||
)
|
||||
quota_commit.assert_called_once_with(
|
||||
tenant_id=tenant_id,
|
||||
@@ -389,7 +404,11 @@ def test_check_and_deduct_credits_uses_billing_reserve_and_commit_when_enabled()
|
||||
bucket="trial",
|
||||
reservation_id="reservation-1",
|
||||
actual_amount=3,
|
||||
meta={"source": "credit_pool.check_and_deduct"},
|
||||
meta={
|
||||
"source": "credit_pool.check_and_deduct",
|
||||
"created_by": CreditUsageCreatedBy.UNKNOWN.value,
|
||||
"app_type": CreditUsageAppType.UNKNOWN.value,
|
||||
},
|
||||
)
|
||||
quota_release.assert_not_called()
|
||||
|
||||
@@ -413,6 +432,8 @@ def test_check_and_deduct_credits_forwards_deterministic_billing_identity() -> N
|
||||
assert result == 3
|
||||
expected_metadata = {
|
||||
"source": "credit_pool.check_and_deduct",
|
||||
"created_by": CreditUsageCreatedBy.UNKNOWN.value,
|
||||
"app_type": CreditUsageAppType.UNKNOWN.value,
|
||||
"agent_run_id": "run-1",
|
||||
}
|
||||
quota_reserve.assert_called_once_with(
|
||||
@@ -509,6 +530,13 @@ def test_deduct_credits_capped_uses_billing_consume_capped_when_enabled() -> Non
|
||||
tenant_id=tenant_id,
|
||||
credits_required=5,
|
||||
pool_type=ProviderQuotaType.PAID,
|
||||
request_id="message-1",
|
||||
metadata={
|
||||
"provider": "openai",
|
||||
"model": "gpt-4o",
|
||||
"app_type": CreditUsageAppType.CHATBOT,
|
||||
"created_by": CreditUsageCreatedBy.APP,
|
||||
},
|
||||
)
|
||||
|
||||
assert result == 2
|
||||
@@ -516,9 +544,15 @@ def test_deduct_credits_capped_uses_billing_consume_capped_when_enabled() -> Non
|
||||
tenant_id=tenant_id,
|
||||
feature_key=FEATURE_KEY_CREDIT_POOL,
|
||||
bucket="paid",
|
||||
request_id=ANY,
|
||||
request_id="message-1",
|
||||
amount=5,
|
||||
meta={"source": "credit_pool.deduct_capped"},
|
||||
meta={
|
||||
"source": "credit_pool.deduct_capped",
|
||||
"provider": "openai",
|
||||
"model": "gpt-4o",
|
||||
"app_type": CreditUsageAppType.CHATBOT.value,
|
||||
"created_by": CreditUsageCreatedBy.APP.value,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -10,6 +10,7 @@ stay fast and focus on the wiring itself.
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from core.app.app_config.entities import ModelConfig
|
||||
from core.credit_usage import CreditUsageAppType, CreditUsageCreatedBy
|
||||
from graphon.model_runtime.entities.llm_entities import LLMMode
|
||||
from services.workflow_generator_service import WorkflowGeneratorService
|
||||
|
||||
@@ -63,7 +64,13 @@ class TestWorkflowGeneratorService:
|
||||
)
|
||||
|
||||
# Assert
|
||||
mock_model_manager.for_tenant.assert_called_once_with(tenant_id="t-1")
|
||||
mock_model_manager.for_tenant.assert_called_once_with(
|
||||
tenant_id="t-1",
|
||||
request_metadata={
|
||||
"app_type": CreditUsageAppType.WORKFLOW,
|
||||
"created_by": CreditUsageCreatedBy.WORKFLOW_GENERATION,
|
||||
},
|
||||
)
|
||||
mock_workflow_generator.generate_workflow_graph.assert_called_once()
|
||||
call_kwargs = mock_workflow_generator.generate_workflow_graph.call_args.kwargs
|
||||
assert call_kwargs["model_instance"] is instance
|
||||
|
||||
Reference in New Issue
Block a user