diff --git a/python/dify_plugin/__init__.py b/python/dify_plugin/__init__.py index be4022a2..293af9bf 100644 --- a/python/dify_plugin/__init__.py +++ b/python/dify_plugin/__init__.py @@ -1,5 +1,11 @@ -from .plugin import Plugin -from .model.model import ModelProvider -from .tool.tool import ToolProvider +from gevent import monkey -__all__ = ['Plugin', 'ModelProvider', 'ToolProvider'] \ No newline at end of file +# patch all the blocking calls +monkey.patch_all(sys=True) + +from .plugin import Plugin # noqa +from .model.model import ModelProvider # noqa +from .tool.tool import ToolProvider # noqa +from .config.config import DifyPluginEnv # noqa + +__all__ = ['Plugin', 'ModelProvider', 'ToolProvider', 'DifyPluginEnv'] diff --git a/python/dify_plugin/core/runtime/entities/plugin/io.py b/python/dify_plugin/core/runtime/entities/plugin/io.py index 1c2cf769..daabb519 100644 --- a/python/dify_plugin/core/runtime/entities/plugin/io.py +++ b/python/dify_plugin/core/runtime/entities/plugin/io.py @@ -1,6 +1,7 @@ from enum import Enum -from dify_plugin.core.server.base.response_writer import ResponseWriter +from dify_plugin.core.server.__base.request_reader import RequestReader +from dify_plugin.core.server.__base.response_writer import ResponseWriter class PluginInStream: @@ -16,9 +17,15 @@ class PluginInStream: raise ValueError(f"Invalid value for PluginInStream.Event: {v}") def __init__( - self, session_id: str, event: Event, data: dict, writer: ResponseWriter + self, + session_id: str, + event: Event, + data: dict, + reader: RequestReader, + writer: ResponseWriter, ): self.session_id = session_id self.event = event self.data = data + self.reader = reader self.writer = writer diff --git a/python/dify_plugin/core/runtime/request.py b/python/dify_plugin/core/runtime/request.py index ce430367..95da77a5 100644 --- a/python/dify_plugin/core/runtime/request.py +++ b/python/dify_plugin/core/runtime/request.py @@ -42,21 +42,16 @@ from dify_plugin.core.runtime.entities.plugin.workflow import ( ParameterExtractorNodeData, QuestionClassifierNodeData, ) -from dify_plugin.core.server.__base.request_reader import RequestReader -from dify_plugin.core.server.__base.response_writer import ResponseWriter +from dify_plugin.core.runtime.session import Session from dify_plugin.tool.entities import ToolInvokeMessage class RequestInterface(AbstractRequestInterface): def __init__( self, - response_writer: Optional[ResponseWriter], - request_reader: Optional[RequestReader], - session_id: Optional[str] = None, + session: Optional[Session] = None, ) -> None: - self.response_writer = response_writer - self.request_reader = request_reader - self.session_id = session_id + self.session = session def _generate_backwards_request_id(self): return uuid.uuid4().hex @@ -82,12 +77,12 @@ class RequestInterface(AbstractRequestInterface): data_type: Type[T], data: dict, ) -> Generator[T, None, None]: - if not self.response_writer or not self.request_reader: + if not self.session: raise Exception("current tool runtime does not support backwards invoke") - self.response_writer.session_message( - session_id=self.session_id, - data=self.response_writer.stream_invoke_object( + self.session.writer.session_message( + session_id=self.session.session_id, + data=self.session.writer.stream_invoke_object( data={ "type": type.value, "backwards_request_id": backwards_request_id, @@ -103,7 +98,7 @@ class RequestInterface(AbstractRequestInterface): ) empty_response_count = 0 - with self.request_reader.read(filter) as reader: + with self.session.reader.read(filter) as reader: for chunk in reader.read(timeout_for_round=1): """ accept response from input stream and wait for at most 60 seconds diff --git a/python/dify_plugin/core/runtime/session.py b/python/dify_plugin/core/runtime/session.py index 3a368fdb..97b02de0 100644 --- a/python/dify_plugin/core/runtime/session.py +++ b/python/dify_plugin/core/runtime/session.py @@ -1,15 +1,25 @@ from concurrent.futures import ThreadPoolExecutor +from dify_plugin.core.server.__base.request_reader import RequestReader +from dify_plugin.core.server.__base.response_writer import ResponseWriter + class Session: + # class variable to store all sessions _session_pool = set["Session"]() + _executor: ThreadPoolExecutor + # current session id session_id: str - executor: ThreadPoolExecutor + # reader and writer + reader: RequestReader + writer: ResponseWriter - def __init__(self, session_id: str, executor: ThreadPoolExecutor) -> None: + def __init__(self, session_id: str, executor: ThreadPoolExecutor, reader: RequestReader, writer: ResponseWriter): self.session_id = session_id self._session_pool.add(self) - self.executor = executor + self._executor = executor + self.reader = reader + self.writer = writer def __del__(self): self._session_pool.remove(self) diff --git a/python/dify_plugin/core/server/__base/stream_reader.py b/python/dify_plugin/core/server/__base/filter_reader.py similarity index 84% rename from python/dify_plugin/core/server/__base/stream_reader.py rename to python/dify_plugin/core/server/__base/filter_reader.py index 3ebbf28d..8716f4aa 100644 --- a/python/dify_plugin/core/server/__base/stream_reader.py +++ b/python/dify_plugin/core/server/__base/filter_reader.py @@ -1,4 +1,3 @@ -from abc import ABC, abstractmethod from collections.abc import Callable, Generator from queue import Queue import queue @@ -8,7 +7,7 @@ from typing import Optional, overload from dify_plugin.core.runtime.entities.plugin.io import PluginInStream -class PluginReader: +class FilterReader: filter: Callable[[PluginInStream], bool] queue: Queue[PluginInStream | None] close_callback: Optional[Callable] @@ -59,15 +58,3 @@ class PluginReader: def __exit__(self, exc_type, exc_value, traceback): self.close() - - -class PluginInputStreamReader(ABC): - @abstractmethod - def read( - self, - ) -> Generator[ - PluginInStream, None, None - ]: - """ - read data from the stream infinitely - """ diff --git a/python/dify_plugin/core/server/__base/request_reader.py b/python/dify_plugin/core/server/__base/request_reader.py index 91627304..ada93df3 100644 --- a/python/dify_plugin/core/server/__base/request_reader.py +++ b/python/dify_plugin/core/server/__base/request_reader.py @@ -1,34 +1,36 @@ +from abc import ABC, abstractmethod +from collections.abc import Generator import threading -from typing import Callable, Optional +from typing import Callable from dify_plugin.core.runtime.entities.plugin.io import PluginInStream -from dify_plugin.core.server.__base.stream_reader import ( - PluginInputStreamReader, - PluginReader, +from dify_plugin.core.server.__base.filter_reader import ( + FilterReader, ) -class RequestReader: - lock = threading.Lock() - readers: list[PluginReader] = [] - stream_reader: Optional[PluginInputStreamReader] - def __init__(self, reader: PluginInputStreamReader): - self.stream_reader = reader +class RequestReader(ABC): + lock: threading.Lock = threading.Lock() + readers: list[FilterReader] = [] + + @abstractmethod + def _read_stream(self) -> Generator[PluginInStream, None, None]: + """ + Read stream from stdin + """ + raise NotImplementedError def event_loop(self): # read line by line while True: - if self.stream_reader is None: - continue - - for line in self.stream_reader.read(): + for line in self._read_stream(): self._process_line(line) def _process_line(self, data: PluginInStream): try: session_id = data.session_id - readers: list[PluginReader] = [] + readers: list[FilterReader] = [] with self.lock: for reader in self.readers: if reader.filter(data): @@ -43,12 +45,12 @@ class RequestReader: }, ) - def read(self, filter: Callable[[PluginInStream], bool]) -> PluginReader: - def close(reader: PluginReader): + def read(self, filter: Callable[[PluginInStream], bool]) -> FilterReader: + def close(reader: FilterReader): with self.lock: self.readers.remove(reader) - reader = PluginReader(filter, close_callback=lambda: close(reader)) + reader = FilterReader(filter, close_callback=lambda: close(reader)) with self.lock: self.readers.append(reader) diff --git a/python/dify_plugin/core/server/aws/request_reader.py b/python/dify_plugin/core/server/aws/request_reader.py index d6353d8d..334c7e97 100644 --- a/python/dify_plugin/core/server/aws/request_reader.py +++ b/python/dify_plugin/core/server/aws/request_reader.py @@ -3,11 +3,11 @@ import threading from typing import Generator from flask import Flask, request from dify_plugin.core.runtime.entities.plugin.io import PluginInStream +from dify_plugin.core.server.__base.request_reader import RequestReader from dify_plugin.core.server.aws.response_writer import AWSResponseWriter -from dify_plugin.core.server.__base.stream_reader import PluginInputStreamReader -class AWSLambdaRequestReader(PluginInputStreamReader): +class AWSLambdaRequestReader(RequestReader): def __init__(self, port: int, max_single_connection_lifetime: int): """ Initialize the AWSLambdaStream and wait for jobs @@ -20,9 +20,7 @@ class AWSLambdaRequestReader(PluginInputStreamReader): # setup server self._serve() - def read( - self, - ) -> Generator[PluginInStream, None, None]: + def _read_stream(self) -> Generator[PluginInStream, None, None]: """ Read request from http server """ @@ -44,6 +42,7 @@ class AWSLambdaRequestReader(PluginInputStreamReader): event=event, session_id=data["session_id"], data=data["data"], + reader=self, writer=AWSResponseWriter(queue), ) # put request to queue diff --git a/python/dify_plugin/core/server/io_server.py b/python/dify_plugin/core/server/io_server.py index 11686040..b8f60af1 100644 --- a/python/dify_plugin/core/server/io_server.py +++ b/python/dify_plugin/core/server/io_server.py @@ -20,7 +20,9 @@ class IOServer(ABC): self.io_stream.close() @abstractmethod - def _execute_request(self, session_id: str, data: dict): + def _execute_request( + self, session_id: str, data: dict, reader: RequestReader, writer: ResponseWriter + ): """ accept requests and execute them, should be implemented outside """ @@ -40,18 +42,19 @@ class IOServer(ABC): self._execute_request_thread, data.session_id, data.data, + data.reader, data.writer, ) def _execute_request_thread( - self, session_id: str, data: dict, writer: ResponseWriter + self, session_id: str, data: dict, reader: RequestReader, writer: ResponseWriter ): """ wrapper for _execute_request """ # wait for the task to finish try: - self._execute_request(session_id, data) + self._execute_request(session_id, data, reader, writer) except Exception as e: writer.session_message( session_id=session_id, diff --git a/python/dify_plugin/core/server/stdio/request_reader.py b/python/dify_plugin/core/server/stdio/request_reader.py index 9a179d3c..a594801a 100644 --- a/python/dify_plugin/core/server/stdio/request_reader.py +++ b/python/dify_plugin/core/server/stdio/request_reader.py @@ -2,15 +2,15 @@ from json import loads import sys from typing import Generator from dify_plugin.core.runtime.entities.plugin.io import PluginInStream -from dify_plugin.core.server.__base.stream_reader import PluginInputStreamReader from gevent.os import tp_read +from dify_plugin.core.server.__base.request_reader import RequestReader from dify_plugin.core.server.stdio.response_writer import StdioResponseWriter -class StdioRequestReader(PluginInputStreamReader): - def read(self) -> Generator[PluginInStream, None, None]: +class StdioRequestReader(RequestReader): + def _read_stream(self) -> Generator[PluginInStream, None, None]: buffer = "" while True: # read data from stdin through tp_read @@ -38,6 +38,7 @@ class StdioRequestReader(PluginInputStreamReader): session_id=data["session_id"], event=PluginInStream.Event.value_of(data["event"]), data=data["data"], + reader=self, writer=StdioResponseWriter(), ) except Exception as e: diff --git a/python/dify_plugin/core/server/tcp/request_reader.py b/python/dify_plugin/core/server/tcp/request_reader.py index 79bece1a..c8981a93 100644 --- a/python/dify_plugin/core/server/tcp/request_reader.py +++ b/python/dify_plugin/core/server/tcp/request_reader.py @@ -6,15 +6,15 @@ from typing import Callable, Generator, Optional from gevent.select import select from dify_plugin.core.runtime.entities.plugin.io import PluginInStream +from dify_plugin.core.server.__base.request_reader import RequestReader from dify_plugin.core.server.__base.response_writer import ResponseWriter -from dify_plugin.core.server.__base.stream_reader import PluginInputStreamReader import logging logger = logging.getLogger(__name__) -class TCPReaderWriter(PluginInputStreamReader, ResponseWriter): +class TCPReaderWriter(RequestReader, ResponseWriter): def __init__( self, host: str, @@ -94,7 +94,7 @@ class TCPReaderWriter(PluginInputStreamReader, ResponseWriter): logger.error(f"Failed to connect to {self.host}:{self.port}, {e}") raise e - def read(self) -> Generator[PluginInStream, None, None]: + def _read_stream(self) -> Generator[PluginInStream, None, None]: """ Read data from the target """ @@ -137,6 +137,7 @@ class TCPReaderWriter(PluginInputStreamReader, ResponseWriter): session_id=data["session_id"], event=PluginInStream.Event.value_of(data["event"]), data=data["data"], + reader=self, writer=self, ) except Exception: diff --git a/python/dify_plugin/plugin.py b/python/dify_plugin/plugin.py index 80c8d4b9..5bafd8e5 100644 --- a/python/dify_plugin/plugin.py +++ b/python/dify_plugin/plugin.py @@ -1,22 +1,14 @@ from typing import Optional -from gevent import monkey from dify_plugin.core.server.__base.request_reader import RequestReader from dify_plugin.core.server.__base.response_writer import ResponseWriter -from dify_plugin.core.server.__base.stream_reader import PluginInputStreamReader from dify_plugin.core.server.aws.request_reader import AWSLambdaRequestReader from dify_plugin.core.server.stdio.request_reader import StdioRequestReader from dify_plugin.core.server.stdio.response_writer import StdioResponseWriter from dify_plugin.core.server.tcp.request_reader import TCPReaderWriter - -# patch all the blocking calls -monkey.patch_all(sys=True) - from collections.abc import Generator # noqa: E402 import logging # noqa: E402 - from dify_plugin.config.config import DifyPluginEnv, InstallMethod # noqa: E402 - from dify_plugin.core.runtime.entities.plugin.request import ( # noqa: E402 ModelActions, PluginInvokeType, @@ -46,11 +38,11 @@ class Plugin(IOServer, Router): self.registration = PluginRegistration(config) if config.INSTALL_METHOD == InstallMethod.Local: - stream_reader, response_writer = self._launch_local_stream(config) + request_reader, response_writer = self._launch_local_stream(config) elif config.INSTALL_METHOD == InstallMethod.Remote: - stream_reader, response_writer = self._launch_remote_stream(config) + request_reader, response_writer = self._launch_remote_stream(config) elif config.INSTALL_METHOD == InstallMethod.AWSLambda: - stream_reader, response_writer = self._launch_aws_stream(config) + request_reader, response_writer = self._launch_aws_stream(config) else: raise ValueError("Invalid install method") @@ -60,7 +52,6 @@ class Plugin(IOServer, Router): # initialize plugin executor self.plugin_executer = PluginExecutor(self.registration) - request_reader = RequestReader(stream_reader) IOServer.__init__(self, config, request_reader) Router.__init__(self, request_reader, response_writer) @@ -69,7 +60,7 @@ class Plugin(IOServer, Router): def _launch_local_stream( self, config: DifyPluginEnv - ) -> tuple[PluginInputStreamReader, Optional[ResponseWriter]]: + ) -> tuple[RequestReader, Optional[ResponseWriter]]: """ Launch local stream """ @@ -82,7 +73,7 @@ class Plugin(IOServer, Router): def _launch_remote_stream( self, config: DifyPluginEnv - ) -> tuple[PluginInputStreamReader, Optional[ResponseWriter]]: + ) -> tuple[RequestReader, Optional[ResponseWriter]]: """ Launch remote stream """ @@ -105,7 +96,7 @@ class Plugin(IOServer, Router): def _launch_aws_stream( self, config: DifyPluginEnv - ) -> tuple[PluginInputStreamReader, Optional[ResponseWriter]]: + ) -> tuple[RequestReader, Optional[ResponseWriter]]: """ Launch AWS stream """ @@ -195,14 +186,18 @@ class Plugin(IOServer, Router): and data.get("action") == WebhookActions.InvokeWebhook.value, ) - def _execute_request(self, session_id: str, data: dict): + def _execute_request( + self, session_id: str, data: dict, reader: RequestReader, writer: ResponseWriter + ): """ accept requests and execute :param session_id: session id, unique for each request :param data: request data """ - session = Session(session_id=session_id, executor=self.executer) + session = Session( + session_id=session_id, executor=self.executer, reader=reader, writer=writer + ) response = self.dispatch(session, data) if response: if isinstance(response, Generator): diff --git a/python/dify_plugin/tool/entities.py b/python/dify_plugin/tool/entities.py index 9db354cb..621730e3 100644 --- a/python/dify_plugin/tool/entities.py +++ b/python/dify_plugin/tool/entities.py @@ -2,6 +2,7 @@ from typing import Any, Optional, Union from pydantic import BaseModel, Field, field_validator, model_validator from enum import Enum +from dify_plugin.config.config import InstallMethod from dify_plugin.core.runtime.entities.plugin.common import I18nObject from dify_plugin.utils.yaml_loader import load_yaml_file @@ -9,6 +10,7 @@ class ToolRuntime(BaseModel): credentials: dict[str, str] user_id: Optional[str] session_id: Optional[str] + install_method: Optional[InstallMethod] = Field(default=InstallMethod.Local) class ToolInvokeMessage(BaseModel): class TextMessage(BaseModel): diff --git a/python/dify_plugin/tool/tool.py b/python/dify_plugin/tool/tool.py index 59f91228..574dadf8 100644 --- a/python/dify_plugin/tool/tool.py +++ b/python/dify_plugin/tool/tool.py @@ -3,8 +3,7 @@ from collections.abc import Generator from typing import Optional from dify_plugin.core.runtime.request import RequestInterface -from dify_plugin.core.server.__base.request_reader import RequestReader -from dify_plugin.core.server.__base.response_writer import ResponseWriter +from dify_plugin.core.runtime.session import Session from dify_plugin.tool.entities import ToolInvokeMessage, ToolRuntime @@ -23,25 +22,18 @@ class Tool(RequestInterface, ABC): def __init__( self, runtime: ToolRuntime, - request_reader: Optional[RequestReader], - response_writer: Optional[ResponseWriter], + session: Optional[Session] = None, ): self.runtime = runtime - RequestInterface.__init__( - self, response_writer, request_reader, runtime.session_id - ) + RequestInterface.__init__(self, session) @classmethod def from_credentials( cls, credentials: dict, - request_reader: Optional[RequestReader], - response_writer: Optional[ResponseWriter], ) -> "Tool": return cls( ToolRuntime(credentials=credentials, user_id=None, session_id=None), - request_reader=request_reader, - response_writer=response_writer, ) def create_text_message(self, text: str) -> ToolInvokeMessage: diff --git a/python/examples/code_based_workflow/main.py b/python/examples/code_based_workflow/main.py index 978d89ba..cc06b3e6 100644 --- a/python/examples/code_based_workflow/main.py +++ b/python/examples/code_based_workflow/main.py @@ -4,8 +4,7 @@ import sys sys.path.append('../..') -from dify_plugin.config.config import DifyPluginEnv -from dify_plugin.plugin import Plugin +from dify_plugin import Plugin, DifyPluginEnv plugin = Plugin(DifyPluginEnv(MAX_REQUEST_TIMEOUT=30)) diff --git a/python/examples/google/main.py b/python/examples/google/main.py index 978d89ba..cc06b3e6 100644 --- a/python/examples/google/main.py +++ b/python/examples/google/main.py @@ -4,8 +4,7 @@ import sys sys.path.append('../..') -from dify_plugin.config.config import DifyPluginEnv -from dify_plugin.plugin import Plugin +from dify_plugin import Plugin, DifyPluginEnv plugin = Plugin(DifyPluginEnv(MAX_REQUEST_TIMEOUT=30)) diff --git a/python/examples/jina/main.py b/python/examples/jina/main.py index 978d89ba..cc06b3e6 100644 --- a/python/examples/jina/main.py +++ b/python/examples/jina/main.py @@ -4,8 +4,7 @@ import sys sys.path.append('../..') -from dify_plugin.config.config import DifyPluginEnv -from dify_plugin.plugin import Plugin +from dify_plugin import Plugin, DifyPluginEnv plugin = Plugin(DifyPluginEnv(MAX_REQUEST_TIMEOUT=30)) diff --git a/python/examples/neko/main.py b/python/examples/neko/main.py index 8cdafbff..aa0d026a 100644 --- a/python/examples/neko/main.py +++ b/python/examples/neko/main.py @@ -4,8 +4,7 @@ import sys sys.path.append('../..') -from dify_plugin.config.config import DifyPluginEnv -from dify_plugin.plugin import Plugin +from dify_plugin import Plugin, DifyPluginEnv plugin = Plugin(DifyPluginEnv(MAX_REQUEST_TIMEOUT=30))