mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
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:
co-authored by
autofix-ci[bot]
WH-2099
parent
19838972dc
commit
f0b34bdeb4
@@ -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__:
|
||||
|
||||
@@ -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"
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -26,6 +26,7 @@ from core.workflow.node_runtime import (
|
||||
DifyFileReferenceFactory,
|
||||
DifyHumanInputNodeRuntime,
|
||||
DifyPreparedLLM,
|
||||
DifyPreparedPollingLLM,
|
||||
DifyPromptMessageSerializer,
|
||||
DifyRetrieverAttachmentLoader,
|
||||
DifyToolFileManager,
|
||||
@@ -531,7 +532,11 @@ class DifyNodeFactory(NodeFactory):
|
||||
node_init_kwargs: dict[str, object] = {
|
||||
"credentials_provider": self._llm_credentials_provider,
|
||||
"model_factory": self._llm_model_factory,
|
||||
"model_instance": DifyPreparedLLM(model_instance) if wrap_model_instance else model_instance,
|
||||
"model_instance": (
|
||||
self._wrap_model_instance_for_node(node_data=validated_node_data, model_instance=model_instance)
|
||||
if wrap_model_instance
|
||||
else model_instance
|
||||
),
|
||||
"memory": self._build_memory_for_llm_node(
|
||||
node_data=validated_node_data,
|
||||
model_instance=model_instance,
|
||||
@@ -555,6 +560,23 @@ class DifyNodeFactory(NodeFactory):
|
||||
node_init_kwargs["default_query_selector"] = system_variable_selector(SystemVariableKey.QUERY)
|
||||
return node_init_kwargs
|
||||
|
||||
@staticmethod
|
||||
def _wrap_model_instance_for_node(
|
||||
*,
|
||||
node_data: LLMCompatibleNodeData,
|
||||
model_instance: ModelInstance,
|
||||
) -> 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)
|
||||
|
||||
@staticmethod
|
||||
def _supports_plugin_llm_polling(model_instance: ModelInstance) -> bool:
|
||||
return model_instance.get_model_schema().support_polling
|
||||
|
||||
def _build_retriever_attachment_loader(self, node_data: LLMNodeData) -> DifyRetrieverAttachmentLoader:
|
||||
return DifyRetrieverAttachmentLoader(
|
||||
file_reference_factory=self._file_reference_factory,
|
||||
|
||||
@@ -4,6 +4,7 @@ from collections.abc import Callable, Generator, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any, Literal, cast, overload, override
|
||||
|
||||
from pydantic import JsonValue
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
@@ -38,6 +39,7 @@ from factories import file_factory
|
||||
from graphon.file import File, FileTransferMethod, FileType
|
||||
from graphon.model_runtime.entities import LLMMode
|
||||
from graphon.model_runtime.entities.llm_entities import (
|
||||
LLMPollingResult,
|
||||
LLMResult,
|
||||
LLMResultChunk,
|
||||
LLMResultChunkWithStructuredOutput,
|
||||
@@ -54,6 +56,7 @@ from graphon.nodes.human_input.entities import (
|
||||
HumanInputNodeData,
|
||||
)
|
||||
from graphon.nodes.llm.runtime_protocols import (
|
||||
LLMPollingCapableProtocol,
|
||||
LLMProtocol,
|
||||
PromptMessageSerializerProtocol,
|
||||
RetrieverAttachmentLoaderProtocol,
|
||||
@@ -278,6 +281,58 @@ class DifyPreparedLLM(LLMProtocol):
|
||||
return isinstance(error, OutputParserError)
|
||||
|
||||
|
||||
class DifyPreparedPollingLLM(DifyPreparedLLM, LLMPollingCapableProtocol):
|
||||
"""Prepared workflow LLM adapter that exposes Graphon's polling protocol."""
|
||||
|
||||
def __init__(self, model_instance: ModelInstance) -> None:
|
||||
from core.plugin.impl.model_runtime import PluginModelRuntime
|
||||
|
||||
super().__init__(model_instance)
|
||||
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.")
|
||||
|
||||
plugin_model_runtime = model_type_instance.model_runtime
|
||||
if not isinstance(plugin_model_runtime, PluginModelRuntime):
|
||||
raise TypeError("Polling wrapper requires a plugin-backed model runtime.")
|
||||
|
||||
self._plugin_model_runtime = plugin_model_runtime
|
||||
|
||||
@override
|
||||
def start_llm_polling(
|
||||
self,
|
||||
*,
|
||||
prompt_messages: Sequence[PromptMessage],
|
||||
model_parameters: Mapping[str, Any],
|
||||
tools: Sequence[PromptMessageTool] | None,
|
||||
stop: Sequence[str] | None,
|
||||
json_schema: Mapping[str, Any] | None,
|
||||
) -> LLMPollingResult:
|
||||
return self._plugin_model_runtime.start_llm_polling(
|
||||
provider=self.provider,
|
||||
model=self.model_name,
|
||||
credentials=self._model_instance.credentials,
|
||||
prompt_messages=prompt_messages,
|
||||
model_parameters=dict(model_parameters),
|
||||
tools=tools,
|
||||
stop=stop,
|
||||
json_schema=dict(json_schema) if json_schema is not None else None,
|
||||
)
|
||||
|
||||
@override
|
||||
def check_llm_polling(
|
||||
self,
|
||||
*,
|
||||
plugin_state: Mapping[str, JsonValue],
|
||||
) -> LLMPollingResult:
|
||||
return self._plugin_model_runtime.check_llm_polling(
|
||||
provider=self.provider,
|
||||
model=self.model_name,
|
||||
credentials=self._model_instance.credentials,
|
||||
plugin_state=dict(plugin_state),
|
||||
)
|
||||
|
||||
|
||||
class DifyPromptMessageSerializer(PromptMessageSerializerProtocol):
|
||||
@override
|
||||
def serialize(
|
||||
|
||||
Reference in New Issue
Block a user