feat(api): pass app_id to model plugins for provider-side cost attribution (#35859)

Co-authored-by: WH-2099 <wh2099@pm.me>
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
Ryuta KOBAYASHI 2026-07-06 17:54:16 +09:00 committed by GitHub
parent 78ca1a8a76
commit 64d72c6fb9
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
18 changed files with 146 additions and 27 deletions

View File

@ -133,6 +133,7 @@ class CotAgentRunner(BaseAgentRunner, ABC):
stop=app_generate_entity.model_conf.stop,
stream=True,
callbacks=[],
request_metadata={"app_id": self.app_config.app_id},
)
usage_dict: dict[str, LLMUsage | None] = {}

View File

@ -101,6 +101,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},
)
tool_calls: list[tuple[str, str, dict[str, Any]]] = []

View File

@ -232,6 +232,7 @@ class ChatAppRunner(AppRunner):
model_parameters=application_generate_entity.model_conf.parameters,
stop=stop,
stream=application_generate_entity.stream,
request_metadata={"app_id": app_config.app_id},
)
# handle invoke result

View File

@ -193,6 +193,7 @@ class CompletionAppRunner(AppRunner):
model_parameters=application_generate_entity.model_conf.parameters,
stop=stop,
stream=application_generate_entity.stream,
request_metadata={"app_id": app_config.app_id},
)
# handle invoke result

View File

@ -124,6 +124,7 @@ class ModelInstance:
stop: list[str] | None = None,
stream: Literal[True] = True,
callbacks: list[Callback] | None = None,
request_metadata: Mapping[str, object] | None = None,
) -> Generator: ...
@overload
@ -135,6 +136,7 @@ class ModelInstance:
stop: list[str] | None = None,
stream: Literal[False] = False,
callbacks: list[Callback] | None = None,
request_metadata: Mapping[str, object] | None = None,
) -> LLMResult: ...
@overload
@ -146,6 +148,7 @@ class ModelInstance:
stop: list[str] | None = None,
stream: bool = True,
callbacks: list[Callback] | None = None,
request_metadata: Mapping[str, object] | None = None,
) -> Union[LLMResult, Generator]: ...
def invoke_llm(
@ -156,6 +159,7 @@ class ModelInstance:
stop: Sequence[str] | None = None,
stream: bool = True,
callbacks: list[Callback] | None = None,
request_metadata: Mapping[str, object] | None = None,
) -> Union[LLMResult, Generator]:
"""
Invoke large language model
@ -166,6 +170,7 @@ class ModelInstance:
:param stop: stop words
:param stream: is stream response
:param callbacks: callbacks
:param request_metadata: optional request metadata
:return: full response or stream response chunk generator result
"""
if not isinstance(self.model_type_instance, LargeLanguageModel):
@ -182,6 +187,7 @@ class ModelInstance:
stop=list(stop) if stop else None,
stream=stream,
callbacks=callbacks,
request_metadata=request_metadata,
),
)

View File

@ -27,10 +27,12 @@ _POLLING_UNSUPPORTED_ERROR_MESSAGE = "does not support polling"
class PluginModelClient(BasePluginClient):
@staticmethod
def _dispatch_payload(*, user_id: str | None, data: dict[str, Any]) -> dict[str, Any]:
def _dispatch_payload(*, user_id: str | None, data: dict[str, Any], app_id: str | None = None) -> dict[str, Any]:
payload: dict[str, Any] = {"data": data}
if user_id is not None:
payload["user_id"] = user_id
if app_id is not None:
payload["app_id"] = app_id
return payload
def fetch_model_providers(self, tenant_id: str) -> Sequence[PluginModelProviderEntity]:
@ -166,6 +168,7 @@ class PluginModelClient(BasePluginClient):
tools: list[PromptMessageTool] | None = None,
stop: list[str] | None = None,
stream: bool = True,
app_id: str | None = None,
) -> Generator[LLMResultChunk, None, None]:
"""
Invoke llm
@ -188,6 +191,7 @@ class PluginModelClient(BasePluginClient):
"stop": stop,
"stream": stream,
},
app_id=app_id,
)
),
headers={

View File

@ -317,21 +317,39 @@ class PluginModelRuntime(ModelRuntime):
stream: bool,
request_metadata: Mapping[str, object] | None = None,
) -> LLMResult | Generator[LLMResultChunk, None, None]:
del request_metadata
app_id = request_metadata.get("app_id") if request_metadata else None
if not isinstance(app_id, str):
app_id = None
plugin_id, provider_name = self._split_provider(provider)
result = self.client.invoke_llm(
tenant_id=self.tenant_id,
user_id=self.user_id,
plugin_id=plugin_id,
provider=provider_name,
model=model,
credentials=credentials,
model_parameters=model_parameters,
prompt_messages=list(prompt_messages),
tools=tools,
stop=list(stop) if stop else None,
stream=stream,
)
if app_id is None:
result = self.client.invoke_llm(
tenant_id=self.tenant_id,
user_id=self.user_id,
plugin_id=plugin_id,
provider=provider_name,
model=model,
credentials=credentials,
model_parameters=model_parameters,
prompt_messages=list(prompt_messages),
tools=tools,
stop=list(stop) if stop else None,
stream=stream,
)
else:
result = self.client.invoke_llm(
tenant_id=self.tenant_id,
user_id=self.user_id,
plugin_id=plugin_id,
provider=provider_name,
model=model,
credentials=credentials,
model_parameters=model_parameters,
prompt_messages=list(prompt_messages),
tools=tools,
stop=list(stop) if stop else None,
stream=stream,
app_id=app_id,
)
if stream:
return result

View File

@ -549,7 +549,11 @@ class DifyNodeFactory(NodeFactory):
"credentials_provider": self._llm_credentials_provider,
"model_factory": self._llm_model_factory,
"model_instance": (
self._wrap_model_instance_for_node(node_data=validated_node_data, model_instance=model_instance)
self._wrap_model_instance_for_node(
node_data=validated_node_data,
model_instance=model_instance,
request_metadata={"app_id": self._dify_context.app_id},
)
if wrap_model_instance
else model_instance
),
@ -581,13 +585,14 @@ class DifyNodeFactory(NodeFactory):
*,
node_data: LLMCompatibleNodeData,
model_instance: ModelInstance,
request_metadata: Mapping[str, object] | None = None,
) -> DifyPreparedLLM:
# Only graphon's LLM node consumes the polling protocol. Keep classifier
# and extractor nodes on the existing wrapper even if the same model
# advertises polling support.
if node_data.type == BuiltinNodeTypes.LLM and DifyNodeFactory._supports_plugin_llm_polling(model_instance):
return DifyPreparedPollingLLM(model_instance)
return DifyPreparedLLM(model_instance)
return DifyPreparedPollingLLM(model_instance, request_metadata=request_metadata)
return DifyPreparedLLM(model_instance, request_metadata=request_metadata)
@staticmethod
def _supports_plugin_llm_polling(model_instance: ModelInstance) -> bool:

View File

@ -150,8 +150,9 @@ class DifyFileReferenceFactory(FileReferenceFactoryProtocol):
class DifyPreparedLLM(LLMProtocol):
"""Workflow-layer adapter that hides the full `ModelInstance` API from `graphon` nodes."""
def __init__(self, model_instance: ModelInstance) -> None:
def __init__(self, model_instance: ModelInstance, request_metadata: Mapping[str, object] | None = None) -> None:
self._model_instance = model_instance
self._request_metadata = request_metadata
@property
@override
@ -230,6 +231,7 @@ class DifyPreparedLLM(LLMProtocol):
tools=list(tools or []),
stop=list(stop or []),
stream=stream,
request_metadata=self._request_metadata,
)
@overload
@ -283,10 +285,10 @@ class DifyPreparedLLM(LLMProtocol):
class DifyPreparedPollingLLM(DifyPreparedLLM, LLMPollingCapableProtocol):
"""Prepared workflow LLM adapter that exposes Graphon's polling protocol."""
def __init__(self, model_instance: ModelInstance) -> None:
def __init__(self, model_instance: ModelInstance, request_metadata: Mapping[str, object] | None = None) -> None:
from core.plugin.impl.model_runtime import PluginModelRuntime
super().__init__(model_instance)
super().__init__(model_instance, request_metadata=request_metadata)
model_type_instance = model_instance.model_type_instance
if not isinstance(model_type_instance, LargeLanguageModel):
raise TypeError("Polling wrapper requires a large-language-model instance.")

View File

@ -246,5 +246,6 @@ class MockModelClass(PluginModelClient):
tools: list[PromptMessageTool] | None = None,
stop: list[str] | None = None,
stream: bool = True,
app_id: str | None = None,
):
return MockModelClass.mocked_chat_create_stream(model=model, prompt_messages=prompt_messages, tools=tools)

View File

@ -42,6 +42,7 @@ def runner(mocker: MockerFixture):
application_generate_entity.invoke_from = "test"
app_config = MagicMock()
app_config.app_id = "app"
app_config.agent = MagicMock()
app_config.agent.max_iteration = 1
app_config.prompt_template.simple_prompt_template = "Hello {{name}}"
@ -341,6 +342,7 @@ class TestRun:
)
results = list(runner.run(runner.session, message, "query", {}))
assert runner.model_instance.invoke_llm.call_args.kwargs["request_metadata"] == {"app_id": "app"}
assert results[-1].delta.message.content == ""
def test_run_usage_missing_key_branch(self, runner: DummyRunner, mocker: MockerFixture):

View File

@ -81,6 +81,7 @@ def runner(mocker: MockerFixture):
mocker.patch("core.agent.fc_agent_runner.LLMResultChunkDelta", MagicMock)
app_config = MagicMock()
app_config.app_id = "app"
app_config.agent = MagicMock(max_iteration=2)
app_config.prompt_template = MagicMock(simple_prompt_template="system")
@ -299,6 +300,7 @@ class TestRunMethod:
outputs = list(runner.run(runner.session, message, "query"))
assert len(outputs) == 1
assert runner.model_instance.invoke_llm.call_args.kwargs["request_metadata"] == {"app_id": "app"}
runner.queue_manager.publish.assert_called()
queue_calls = runner.queue_manager.publish.call_args_list

View File

@ -156,15 +156,35 @@ class TestPluginModelClient:
tools=[],
stop=["STOP"],
stream=False,
app_id="app-1",
)
)
assert result == ["chunk-1"]
call_kwargs = stream_mock.call_args.kwargs
assert call_kwargs["path"] == "plugin/tenant-1/dispatch/llm/invoke"
assert call_kwargs["data"]["app_id"] == "app-1"
assert call_kwargs["data"]["data"]["stream"] is False
assert call_kwargs["data"]["data"]["model_parameters"] == {"temperature": 0.1}
def test_invoke_llm_omits_app_id_when_missing(self, mocker: MockerFixture):
client = PluginModelClient()
stream_mock = mocker.patch.object(client, "_request_with_plugin_daemon_response_stream", return_value=iter([]))
list(
client.invoke_llm(
tenant_id="tenant-1",
user_id="user-1",
plugin_id="org/plugin:1",
provider="provider-a",
model="gpt-test",
credentials={},
prompt_messages=[],
)
)
assert "app_id" not in stream_mock.call_args.kwargs["data"]
def test_invoke_llm_wraps_plugin_daemon_inner_error(self, mocker: MockerFixture):
client = PluginModelClient()

View File

@ -278,6 +278,59 @@ class TestPluginModelRuntime:
stream=False,
)
def test_invoke_llm_forwards_string_app_id_from_request_metadata(self) -> None:
client = Mock(spec=PluginModelClient)
client.invoke_llm.return_value = iter([])
runtime = PluginModelRuntime(tenant_id="tenant", user_id="user", client=client, plugin_service=PluginService)
result = runtime.invoke_llm(
provider="langgenius/openai/openai",
model="gpt-4o-mini",
credentials={"api_key": "secret"},
model_parameters={"temperature": 0.3},
prompt_messages=[],
tools=None,
stop=None,
stream=True,
request_metadata={"app_id": "app-1"},
)
assert list(result) == []
client.invoke_llm.assert_called_once_with(
tenant_id="tenant",
user_id="user",
plugin_id="langgenius/openai",
provider="openai",
model="gpt-4o-mini",
credentials={"api_key": "secret"},
model_parameters={"temperature": 0.3},
prompt_messages=[],
tools=None,
stop=None,
stream=True,
app_id="app-1",
)
def test_invoke_llm_ignores_non_string_app_id_request_metadata(self) -> None:
client = Mock(spec=PluginModelClient)
client.invoke_llm.return_value = iter([])
runtime = PluginModelRuntime(tenant_id="tenant", user_id="user", client=client, plugin_service=PluginService)
result = runtime.invoke_llm(
provider="langgenius/openai/openai",
model="gpt-4o-mini",
credentials={"api_key": "secret"},
model_parameters={"temperature": 0.3},
prompt_messages=[],
tools=None,
stop=None,
stream=True,
request_metadata={"app_id": 123},
)
assert result is client.invoke_llm.return_value
assert "app_id" not in client.invoke_llm.call_args.kwargs
def test_invoke_llm_returns_plugin_stream_directly(self) -> None:
client = Mock(spec=PluginModelClient)
stream_result = iter([])

View File

@ -723,6 +723,7 @@ class TestDifyNodeFactoryCreateNode:
wrap_model.assert_called_once_with(
node_data=node_data,
model_instance=sentinel.model_instance,
request_metadata={"app_id": "app-id"},
)
assert kwargs["model_instance"] is wrapped_model_instance

View File

@ -171,7 +171,7 @@ def test_dify_prepared_llm_wraps_model_instance_calls() -> None:
model_schema = _build_model_schema()
model_instance = _ModelInstanceStub(model_schema=model_schema)
model_type_instance = model_instance.model_type_instance
prepared = DifyPreparedLLM(model_instance)
prepared = DifyPreparedLLM(model_instance, request_metadata={"app_id": "app-id"})
assert prepared.provider == "langgenius/openai/openai"
assert prepared.model_name == "gpt-4o-mini"
@ -197,6 +197,7 @@ def test_dify_prepared_llm_wraps_model_instance_calls() -> None:
tools=[],
stop=[],
stream=False,
request_metadata={"app_id": "app-id"},
)

View File

@ -86,7 +86,7 @@ export type ProviderConfig = {
placeholder?: I18nObject | null
required?: boolean
scope?: AppSelectorScope | ModelSelectorScope | ToolSelectorScope | null
type: CoreEntitiesProviderEntitiesBasicProviderConfigType
type: ProviderConfigType
url?: string | null
}
@ -126,7 +126,7 @@ export type ModelSelectorScope
export type ToolSelectorScope = 'all' | 'builtin' | 'custom' | 'workflow'
export type CoreEntitiesProviderEntitiesBasicProviderConfigType
export type ProviderConfigType
= | 'app-selector'
| 'array[tools]'
| 'boolean'

View File

@ -119,9 +119,9 @@ export const zModelSelectorScope = z.enum([
export const zToolSelectorScope = z.enum(['all', 'builtin', 'custom', 'workflow'])
/**
* Type
* ProviderConfigType
*/
export const zCoreEntitiesProviderEntitiesBasicProviderConfigType = z.enum([
export const zProviderConfigType = z.enum([
'app-selector',
'array[tools]',
'boolean',
@ -146,7 +146,7 @@ export const zProviderConfig = z.object({
placeholder: zI18nObject.nullish(),
required: z.boolean().optional().default(false),
scope: z.union([zAppSelectorScope, zModelSelectorScope, zToolSelectorScope]).nullish(),
type: zCoreEntitiesProviderEntitiesBasicProviderConfigType,
type: zProviderConfigType,
url: z.string().nullish(),
})