feat(api): LLM polling support (#37462)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: WH-2099 <wh2099@pm.me>
This commit is contained in:
QuantumGhost
2026-06-17 23:34:33 +00:00
committed by GitHub
co-authored by autofix-ci[bot] WH-2099
parent 19838972dc
commit f0b34bdeb4
17 changed files with 704 additions and 46 deletions
+5
View File
@@ -20,6 +20,7 @@ from core.plugin.impl.exc import (
PluginDaemonNotFoundError,
PluginDaemonUnauthorizedError,
PluginInvokeError,
PluginLLMPollingUnsupportedError,
PluginNotFoundError,
PluginPermissionDeniedError,
PluginUniqueIdentifierError,
@@ -370,6 +371,10 @@ class BasePluginClient:
raise TriggerInvokeError(error_object.get("message"))
case EventIgnoreError.__name__:
raise EventIgnoreError(description=error_object.get("message"))
# NOTE: current plugin sdk / plugin daemon does not raise exception with
# type `PluginLLMPollingUnsupportedError`.
case PluginLLMPollingUnsupportedError.__name__:
raise PluginLLMPollingUnsupportedError(description=error_object.get("message"))
case _:
raise PluginInvokeError(description=message)
case PluginDaemonInternalServerError.__name__:
+11
View File
@@ -5,6 +5,13 @@ from pydantic import TypeAdapter
from extensions.ext_logging import get_request_id
# NOTE: Avoid renaming exception classes in this file, since
# the `_handle_plugin_daemon_error` in api/core/plugin/impl/base.py
# build exception instances based on the class name.
#
# Renaming of exception classes could result in incorrect exception
# being raised.
class PluginDaemonError(Exception):
"""Base class for all plugin daemon errors."""
@@ -75,6 +82,10 @@ class PluginInvokeError(PluginDaemonClientSideError, ValueError):
)
class PluginLLMPollingUnsupportedError(PluginInvokeError):
"""Plugin-backed LLM polling is unavailable for the requested model."""
class PluginUniqueIdentifierError(PluginDaemonClientSideError):
description: str = "Unique Identifier Error"
+103 -2
View File
@@ -13,13 +13,17 @@ from core.plugin.entities.plugin_daemon import (
PluginVoicesResponse,
)
from core.plugin.impl.base import BasePluginClient
from graphon.model_runtime.entities.llm_entities import LLMResultChunk
from core.plugin.impl.exc import PluginInvokeError, PluginLLMPollingUnsupportedError
from graphon.model_runtime.entities.llm_entities import LLMPollingResult, LLMResultChunk
from graphon.model_runtime.entities.message_entities import PromptMessage, PromptMessageTool
from graphon.model_runtime.entities.model_entities import AIModelEntity
from graphon.model_runtime.entities.model_entities import AIModelEntity, ModelType
from graphon.model_runtime.entities.rerank_entities import MultimodalRerankInput, RerankResult
from graphon.model_runtime.entities.text_embedding_entities import EmbeddingResult
from graphon.model_runtime.utils.encoders import jsonable_encoder
_POLLING_UNSUPPORTED_INVOKE_ERROR_TYPES = frozenset((NotImplementedError.__name__,))
_POLLING_UNSUPPORTED_ERROR_MESSAGE = "does not support polling"
class PluginModelClient(BasePluginClient):
@staticmethod
@@ -197,6 +201,103 @@ class PluginModelClient(BasePluginClient):
except PluginDaemonInnerError as e:
raise ValueError(e.message + str(e.code))
def start_llm_polling(
self,
tenant_id: str,
user_id: str | None,
plugin_id: str,
provider: str,
model: str,
credentials: dict[str, Any],
prompt_messages: list[PromptMessage],
model_parameters: dict[str, Any] | None = None,
tools: list[PromptMessageTool] | None = None,
stop: list[str] | None = None,
json_schema: dict[str, Any] | None = None,
) -> LLMPollingResult:
"""Start an LLM polling request for plugin-backed long-running jobs."""
try:
return self._request_with_plugin_daemon_response(
method="POST",
path=f"plugin/{tenant_id}/dispatch/model/polling/start",
type_=LLMPollingResult,
data=jsonable_encoder(
self._dispatch_payload(
user_id=user_id,
data={
"provider": provider,
"model_type": ModelType.LLM.value,
"model": model,
"credentials": credentials,
"prompt_messages": prompt_messages,
"model_parameters": model_parameters,
"tools": tools,
"stop": stop,
"stream": False,
"json_schema": json_schema,
},
)
),
headers={
"X-Plugin-ID": plugin_id,
"Content-Type": "application/json",
},
)
except PluginInvokeError as error:
self._raise_typed_polling_unsupported_error(error)
raise
def check_llm_polling(
self,
tenant_id: str,
user_id: str | None,
plugin_id: str,
provider: str,
model: str,
credentials: dict[str, Any],
plugin_state: dict[str, Any],
) -> LLMPollingResult:
"""Check the latest state for a plugin-backed LLM polling job."""
try:
return self._request_with_plugin_daemon_response(
method="POST",
path=f"plugin/{tenant_id}/dispatch/model/polling/check",
type_=LLMPollingResult,
data=jsonable_encoder(
self._dispatch_payload(
user_id=user_id,
data={
"provider": provider,
"model_type": ModelType.LLM.value,
"model": model,
"credentials": credentials,
"plugin_state": plugin_state,
},
)
),
headers={
"X-Plugin-ID": plugin_id,
"Content-Type": "application/json",
},
)
except PluginInvokeError as error:
self._raise_typed_polling_unsupported_error(error)
raise
@staticmethod
def _raise_typed_polling_unsupported_error(error: PluginInvokeError) -> None:
"""Convert plugin polling capability failures into a dedicated Dify exception."""
if error.get_error_type() == PluginLLMPollingUnsupportedError.__name__:
raise PluginLLMPollingUnsupportedError(description=error.description) from error
if (
error.get_error_type() in _POLLING_UNSUPPORTED_INVOKE_ERROR_TYPES
# This is ugly, we should not rely on error messages while checking
# error types.
and _POLLING_UNSUPPORTED_ERROR_MESSAGE in error.get_error_message().lower()
):
raise PluginLLMPollingUnsupportedError(description=error.description) from error
def get_llm_num_tokens(
self,
tenant_id: str,
+50
View File
@@ -6,6 +6,7 @@ from collections.abc import Generator, Iterable, Sequence
from typing import IO, Any, Literal, cast, overload, override
from pydantic import ValidationError
from pydantic.json_schema import JsonValue
from redis import RedisError
from configs import dify_config
@@ -17,6 +18,7 @@ from core.plugin.impl.model import PluginModelClient
from core.plugin.plugin_service import PluginService
from extensions.ext_redis import redis_client
from graphon.model_runtime.entities.llm_entities import (
LLMPollingResult,
LLMResult,
LLMResultChunk,
LLMResultChunkWithStructuredOutput,
@@ -430,6 +432,54 @@ class PluginModelRuntime(ModelRuntime):
tools=list(tools) if tools else None,
)
def start_llm_polling(
self,
*,
provider: str,
model: str,
credentials: dict[str, Any],
model_parameters: dict[str, Any],
prompt_messages: Sequence[PromptMessage],
tools: Sequence[PromptMessageTool] | None,
stop: Sequence[str] | None,
json_schema: dict[str, Any] | None,
) -> LLMPollingResult:
"""Start a plugin-side polling job for long-running LLM invocations."""
plugin_id, provider_name = self._split_provider(provider)
return self.client.start_llm_polling(
tenant_id=self.tenant_id,
user_id=self.user_id,
plugin_id=plugin_id,
provider=provider_name,
model=model,
credentials=credentials,
prompt_messages=list(prompt_messages),
model_parameters=model_parameters,
tools=list(tools) if tools else None,
stop=list(stop) if stop else None,
json_schema=json_schema,
)
def check_llm_polling(
self,
*,
provider: str,
model: str,
credentials: dict[str, Any],
plugin_state: dict[str, JsonValue],
) -> LLMPollingResult:
"""Check the latest plugin-side polling state for an LLM invocation."""
plugin_id, provider_name = self._split_provider(provider)
return self.client.check_llm_polling(
tenant_id=self.tenant_id,
user_id=self.user_id,
plugin_id=plugin_id,
provider=provider_name,
model=model,
credentials=credentials,
plugin_state=plugin_state,
)
@override
def invoke_text_embedding(
self,