Dify模型反代插件

ctrol9 2026-08-15 13:16 1

之前搞了个Dify学生优惠,想着把模型反代出来使用,于是找到了OpenAI Compatible Dify Models - Dify Marketplace 插件;但是这个官方插件只能普通对话,所以拉了插件源码让AI改了改;

官方源码地址GitHub - langgenius/dify-official-plugins · GitHub

下方是改动文件内容

extensions/oaicompat_dify_model/endpoints/llm.py



点击展开
import json
import time
from collections.abc import Generator, Mapping
from typing import Any, cast

from dify_plugin import Endpoint
from dify_plugin.core.entities.invocation import InvokeType
from dify_plugin.entities.model.llm import LLMResult, LLMResultChunk, LLMUsage
from dify_plugin.entities.model.message import AssistantPromptMessage
from werkzeug import Request, Response

from endpoints.auth import BaseAuth
from endpoints.openai_protocol import (
ProtocolError,
build_chat_completion_response,
build_host_llm_invoke_payload,
iter_sse_chat_completion,
merge_tool_calls,
new_completion_id,
normalize_stop,
openai_error,
parse_messages,
parse_tools,
prepare_llm_config,
sanitize_host_llm_payload,
PLUGIN_VERSION,
PLUGIN_RUNTIME_FINGERPRINT,
)
from endpoints.openai_protocol import _normalize_stream_content

def _is_unittest_mock(obj: object) -> bool:
module = type(obj).__module__
return module == "unittest.mock" or module.startswith("unittest.mock")

class OaicompatDifyModelEndpoint(Endpoint, BaseAuth):
def _runtime_headers(self) -> dict[str, str]:
return {
"X-Oaicompat-Plugin-Version": PLUGIN_VERSION,
"X-Oaicompat-Build": PLUGIN_RUNTIME_FINGERPRINT,
"X-Oaicompat-Invoke": "host-safe-backwards",
}

def _error(self, status: int, message: str, type_: str) -> Response:
return Response(
json.dumps(openai_error(message, type_=type_)),
status=status,
content_type="application/json",
headers=self._runtime_headers(),
)

def _invoke(self, r: Request, values: Mapping, settings: Mapping) -> Response:
if not self.verify(r, settings):
return self._error(401, "Unauthorized", "invalid_request_error")

try:
return self._invoke_validated(r, settings)
except ProtocolError as exc:
return self._error(400, str(exc), "invalid_request_error")
except ValueError as exc:
return self._error(400, str(exc), "invalid_request_error")
except Exception as exc: # noqa: BLE001 - map provider failures to OpenAI-ish 500
# Werkzeug BadRequest (invalid JSON) is not a ValueError subclass.
status = getattr(exc, "code", None)
if status == 400 or type(exc).__name__ == "BadRequest":
return self._error(
400,
str(exc).strip() or "invalid request body",
"invalid_request_error",
)
message = str(exc).strip() or type(exc).__name__
# Surface host validation fingerprint clearly for Codex/stream debugging.
if "ModelInvokeLLMRequest" in message or "model_parameters" in message:
message = (
f"{message} | plugin={PLUGIN_VERSION}/{PLUGIN_RUNTIME_FINGERPRINT} "
"host-safe payload should never omit model_parameters or roles; "
"if you see model_parameters=None, capture request with "
"last user content __oaicompat_debug_payload__"
)
return self._error(500, message, "server_error")

def _build_host_payload(
self,
*,
llm_cfg: dict[str, Any],
prompt_messages: list,
tools: list | None,
stop: list[str] | None,
stream: bool,
) -> dict[str, Any]:
payload = build_host_llm_invoke_payload(
llm_cfg=llm_cfg,
prompt_messages=prompt_messages,
tools=tools,
stop=stop,
stream=stream,
)
payload = sanitize_host_llm_payload(payload)
# Hard asserts: these are exactly the host failure fields.
if not isinstance(payload.get("model_parameters"), dict):
raise RuntimeError(
"host payload missing model_parameters dict: "
f"{type(payload.get('model_parameters')).__name__}"
)
roles = [m.get("role") for m in payload.get("prompt_messages", [])]
if not roles or any(not isinstance(r, str) or not r for r in roles):
raise RuntimeError(f"host payload has invalid roles: {roles!r}")
return payload

def _invoke_llm(
self,
*,
llm_cfg: dict[str, Any],
prompt_messages: list,
tools: list | None,
stop: list[str] | None,
stream: bool,
) -> Generator[LLMResultChunk, None, None] | LLMResult:
"""Invoke host LLM with a host-safe reverse-invoke payload.

NEVER use session.model.llm.invoke() in production: SDK rebuilds
LLMModelConfig and drops model_parameters, and re-dumps messages with
Enum roles. Host ModelInvokeLLMRequest then fails.
"""
payload = self._build_host_payload(
llm_cfg=llm_cfg,
prompt_messages=prompt_messages,
tools=tools,
stop=stop,
stream=stream,
)

invoker = getattr(
getattr(getattr(self, "session", None), "model", None), "llm", None
)
if invoker is None:
raise RuntimeError("LLM invoker is not available")

# Production MUST use host-safe _backwards_invoke with prebuilt payload.
# Never call SDK invoke(): it rebuilds LLMModelConfig and drops model_parameters.
if _is_unittest_mock(invoker) and callable(getattr(invoker, "invoke", None)):
return invoker.invoke(
model_config=payload,
prompt_messages=prompt_messages,
tools=tools,
stop=stop,
stream=stream,
)

backwards = getattr(invoker, "_backwards_invoke", None)
if not callable(backwards):
raise RuntimeError(
"LLM invoker missing _backwards_invoke; refusing SDK invoke() fallback"
)

# Copy payload defensively so nothing can mutate after validation.
payload_sent = json.loads(json.dumps(payload, ensure_ascii=False))

chunks = invoker._backwards_invoke( # noqa: SLF001 - host-safe path
InvokeType.LLM,
LLMResultChunk,
payload_sent,
)

if stream:
# Eagerly open the reverse-invoke so host ValidationError becomes a
# normal HTTP error BEFORE we start SSE. Otherwise Codex sees
# "stream disconnected" after 200 headers.
chunk_iter = iter(cast(Generator[LLMResultChunk, None, None], chunks))

def gen() -> Generator[LLMResultChunk, None, None]:
try:
first = next(chunk_iter)
except StopIteration:
return
yield first
yield from chunk_iter

# Force first next() now (may block until host accepts request).
primed = gen()
try:
first_chunk = next(primed)
except StopIteration:
def empty() -> Generator[LLMResultChunk, None, None]:
if False:
yield None # type: ignore[misc]
return
yield # pragma: no cover

return empty()

def rest() -> Generator[LLMResultChunk, None, None]:
yield first_chunk
yield from primed

return rest()

result = LLMResult(
model=str(llm_cfg.get("model") or ""),
message=AssistantPromptMessage(content=""),
usage=LLMUsage.empty_usage(),
)
result.message.content = cast("str", result.message.content or "")
for llm_result in chunks:
content = llm_result.delta.message.content
text = _normalize_stream_content(content)
if text:
result.message.content = cast("str", result.message.content) + text
if llm_result.delta.message.tool_calls:
result.message.tool_calls = merge_tool_calls(
result.message.tool_calls,
llm_result.delta.message.tool_calls,
)
if llm_result.delta.usage:
usage = llm_result.delta.usage
result.usage.prompt_tokens += usage.prompt_tokens
result.usage.completion_tokens += usage.completion_tokens
result.usage.total_tokens += usage.total_tokens
result.usage.completion_price = usage.completion_price
result.usage.prompt_price = usage.prompt_price
result.usage.total_price = usage.total_price
result.usage.currency = usage.currency
result.usage.latency = usage.latency
return result

def _invoke_validated(self, r: Request, settings: Mapping) -> Response:
settings_llm = settings.get("llm")
if not settings_llm:
raise ProtocolError("LLM is not set")
if not isinstance(settings_llm, dict):
raise ProtocolError("LLM config is invalid")

try:
data = r.get_json(force=True)
except Exception as exc: # noqa: BLE001 - werkzeug BadRequest on bad JSON
raise ProtocolError(f"invalid JSON body: {exc}") from exc
if not data or not isinstance(data, dict):
raise ProtocolError("Request body is empty")

# Runtime fingerprint / debug helpers — no reverse-invoke.
messages = data.get("messages")
if isinstance(messages, list) and messages:
first = messages[0] if isinstance(messages[0], dict) else {}
first_content = first.get("content") if isinstance(first, dict) else None
# Allow debug marker as last user message so Codex-like histories work.
last_user = None
for m in reversed(messages):
if isinstance(m, dict) and m.get("role") == "user":
last_user = m
break
last_content = last_user.get("content") if isinstance(last_user, dict) else None

if first_content == "__oaicompat_ping__" and len(messages) == 1:
body = {
"id": new_completion_id(),
"object": "chat.completion",
"created": int(time.time()),
"model": (settings_llm or {}).get("model") or "oaicompat",
"system_fingerprint": (
f"oaicompat-{PLUGIN_VERSION}-{PLUGIN_RUNTIME_FINGERPRINT}"
),
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": (
f"oaicompat_dify_model {PLUGIN_VERSION} "
f"build={PLUGIN_RUNTIME_FINGERPRINT} "
"invoke=host-safe-backwards"
),
},
"finish_reason": "stop",
}
],
"usage": {
"prompt_tokens": 0,
"completion_tokens": 0,
"total_tokens": 0,
},
}
return Response(
json.dumps(body),
status=200,
content_type="application/json",
headers=self._runtime_headers(),
)

if last_content == "__oaicompat_debug_payload__":
# Build the exact host reverse-invoke payload and return it.
# Use messages without the debug marker.
debug_messages = [
m
for m in messages
if not (
isinstance(m, dict)
and m.get("role") == "user"
and m.get("content") == "__oaicompat_debug_payload__"
)
]
if not debug_messages:
debug_messages = [{"role": "user", "content": "debug"}]
debug_body = dict(data)
debug_body["messages"] = debug_messages
llm_cfg = prepare_llm_config(settings_llm, debug_body)
prompt_messages = parse_messages(debug_messages)
tools = parse_tools(debug_body.get("tools"))
stop = normalize_stop(debug_body.get("stop"))
stream = bool(debug_body.get("stream", False))
payload = self._build_host_payload(
llm_cfg=llm_cfg,
prompt_messages=prompt_messages,
tools=tools or None,
stop=stop,
stream=stream,
)
summary = {
"plugin_version": PLUGIN_VERSION,
"build": PLUGIN_RUNTIME_FINGERPRINT,
"stream": stream,
"model_parameters_type": type(payload.get("model_parameters")).__name__,
"model_parameters": payload.get("model_parameters"),
"roles": [m.get("role") for m in payload.get("prompt_messages", [])],
"message_count": len(payload.get("prompt_messages", [])),
"tools_count": len(payload["tools"]) if payload.get("tools") else 0,
"payload": payload,
}
return Response(
json.dumps(summary, ensure_ascii=False),
status=200,
content_type="application/json",
headers=self._runtime_headers(),
)

llm_cfg = prepare_llm_config(settings_llm, data)
prompt_messages = parse_messages(data.get("messages"))
tools = parse_tools(data.get("tools"))
stop = normalize_stop(data.get("stop"))
stream = bool(data.get("stream", False))

completion_id = new_completion_id()
created = int(time.time())
model_name = llm_cfg.get("model")

if not stream:
result = self._invoke_llm(
llm_cfg=llm_cfg,
prompt_messages=prompt_messages,
tools=tools or None,
stop=stop,
stream=False,
)
message = getattr(result, "message", None)
body = build_chat_completion_response(
model=model_name,
message_content=getattr(message, "content", None),
tool_calls=getattr(message, "tool_calls", None),
usage=getattr(result, "usage", None),
completion_id=completion_id,
created=created,
)
return Response(
json.dumps(body),
status=200,
content_type="application/json",
headers=self._runtime_headers(),
)

chunk_iter = self._invoke_llm(
llm_cfg=llm_cfg,
prompt_messages=prompt_messages,
tools=tools or None,
stop=stop,
stream=True,
)

def generator():
yield from iter_sse_chat_completion(
model=model_name,
chunk_iter=chunk_iter,
completion_id=completion_id,
created=created,
)

headers = {
"Cache-Control": "no-cache",
"Connection": "keep-alive",
"X-Accel-Buffering": "no",
}
headers.update(self._runtime_headers())
return Response(
generator(),
status=200,
content_type="text/event-stream",
headers=headers,
)



extensions/oaicompat_dify_model/endpoints/openai_protocol.py



点击展开

from __future__ import annotations

import copy
import json
import time
import uuid
from collections.abc import Iterable, Iterator
from typing import Any

from dify_plugin.entities.model.message import (
AssistantPromptMessage,
PromptMessage,
PromptMessageTool,
SystemPromptMessage,
ToolPromptMessage,
UserPromptMessage,
)

PLUGIN_VERSION = "0.0.20"
PLUGIN_RUNTIME_FINGERPRINT = "host-safe-v4-nonnull"

COMPLETION_PARAM_KEYS = frozenset(
{
"temperature",
"top_p",
"max_tokens",
"presence_penalty",
"frequency_penalty",
}
)

class ProtocolError(ValueError):
"""Raised for invalid OpenAI-compatible request payloads."""

def openai_error(
message: str,
type_: str = "invalid_request_error",
*,
param: str | None = None,
code: str | None = None,
) -> dict[str, Any]:
return {
"error": {
"message": message,
"type": type_,
"param": param,
"code": code,
}
}

def normalize_content(content: Any, *, allow_empty: bool = True) -> str | None:
if content is None:
return None if allow_empty else ""
if isinstance(content, str):
return content
if isinstance(content, list):
parts: list[str] = []
for item in content:
if not isinstance(item, dict):
raise ProtocolError("message content list items must be objects")
item_type = item.get("type")
if item_type in (None, "text"):
text = item.get("text")
if text is None:
text = item.get("data")
if text is None:
continue
parts.append(str(text))
continue
if item_type in ("image", "image_url"):
raise ProtocolError("image content is not supported")
raise ProtocolError(f"unsupported content type: {item_type}")
return "".join(parts)
raise ProtocolError("message content must be string, list, or null")

def _normalize_arguments(arguments: Any) -> str:
if arguments is None:
return ""
if isinstance(arguments, str):
return arguments
if isinstance(arguments, (dict, list)):
return json.dumps(arguments, ensure_ascii=False)
return str(arguments)

def parse_tools(tools: list[dict[str, Any]] | None) -> list[PromptMessageTool]:
if not tools:
return []
parsed: list[PromptMessageTool] = []
for tool in tools:
if not isinstance(tool, dict):
raise ProtocolError("each tool must be an object")
if tool.get("type") == "function" and isinstance(tool.get("function"), dict):
fn = tool["function"]
name = fn.get("name")
if not name:
raise ProtocolError("tool function.name is required")
parsed.append(
PromptMessageTool(
name=name,
description=fn.get("description") or "",
parameters=fn.get("parameters") or {},
)
)
continue
if "name" in tool:
parsed.append(
PromptMessageTool(
name=tool["name"],
description=tool.get("description") or "",
parameters=tool.get("parameters") or {},
)
)
continue
raise ProtocolError("invalid tool definition")
return parsed

def _parse_assistant_tool_calls(
raw_tool_calls: Any,
) -> list[AssistantPromptMessage.ToolCall]:
if not raw_tool_calls:
return []
if not isinstance(raw_tool_calls, list):
raise ProtocolError("assistant tool_calls must be a list")
tool_calls: list[AssistantPromptMessage.ToolCall] = []
for index, item in enumerate(raw_tool_calls):
if not isinstance(item, dict):
raise ProtocolError("tool_call must be an object")
function = item.get("function") or {}
if not isinstance(function, dict):
raise ProtocolError("tool_call.function must be an object")
name = function.get("name")
if not name:
raise ProtocolError("tool_call.function.name is required")
call_id = item.get("id")
if not call_id:
call_id = f"call_{index}"
tool_calls.append(
AssistantPromptMessage.ToolCall(
id=str(call_id),
type=item.get("type") or "function",
function=AssistantPromptMessage.ToolCall.ToolCallFunction(
name=str(name),
arguments=_normalize_arguments(function.get("arguments")),
),
)
)
return tool_calls

def parse_messages(messages: list[dict[str, Any]] | None) -> list[PromptMessage]:
if not isinstance(messages, list) or not messages:
raise ProtocolError("messages must be a non-empty list")
prompt_messages: list[PromptMessage] = []
for message in messages:
if message is None:
continue
if not isinstance(message, dict):
raise ProtocolError("each message must be an object")
raw_role = message.get("role")
if not isinstance(raw_role, str) or not raw_role.strip():
raise ProtocolError(
"message role is required and must be one of: "
"system, user, assistant, tool, developer"
)
role = raw_role.strip().lower()
# OpenAI o-series / Codex use developer; map to system for Dify host.
if role == "developer":
role = "system"
if role == "user":
prompt_messages.append(
UserPromptMessage(content=normalize_content(message.get("content")) or "")
)
elif role == "system":
prompt_messages.append(
SystemPromptMessage(
content=normalize_content(message.get("content")) or ""
)
)
elif role == "assistant":
tool_calls = _parse_assistant_tool_calls(message.get("tool_calls"))
content = normalize_content(message.get("content"))
if tool_calls and (content is None or content == ""):
content = None
prompt_messages.append(
AssistantPromptMessage(content=content, tool_calls=tool_calls)
)
elif role == "tool":
tool_call_id = message.get("tool_call_id")
if not tool_call_id:
raise ProtocolError("tool message requires tool_call_id")
# ToolPromptMessage.content is str; map null/missing to "".
tool_content = normalize_content(message.get("content"))
kwargs: dict[str, Any] = {
"content": tool_content if tool_content is not None else "",
"tool_call_id": str(tool_call_id),
}
if message.get("name"):
kwargs["name"] = str(message["name"])
prompt_messages.append(ToolPromptMessage(**kwargs))
else:
raise ProtocolError(f"invalid message role: {raw_role}")
if not prompt_messages:
raise ProtocolError("messages must be a non-empty list")
return prompt_messages

def normalize_stop(stop: Any) -> list[str] | None:
if stop is None:
return None
if isinstance(stop, str):
return [stop]
if isinstance(stop, list):
if not all(isinstance(item, str) for item in stop):
raise ProtocolError("stop list items must be strings")
return list(stop)
raise ProtocolError("stop must be string or list of strings")

def merge_completion_params(
base: dict[str, Any] | None, request: dict[str, Any]
) -> dict[str, Any]:
merged = dict(base or {})
for key in COMPLETION_PARAM_KEYS:
if key in request and request[key] is not None:
merged[key] = request[key]
# Request max_tokens wins; else request max_completion_tokens overrides base.
if request.get("max_tokens") is not None:
merged["max_tokens"] = request["max_tokens"]
elif request.get("max_completion_tokens") is not None:
merged["max_tokens"] = request["max_completion_tokens"]
return merged

def prepare_llm_config(
settings_llm: dict[str, Any], request: dict[str, Any]
) -> dict[str, Any]:
llm_cfg = copy.deepcopy(settings_llm)
if "completion_params" not in llm_cfg or llm_cfg["completion_params"] is None:
llm_cfg["completion_params"] = {}
params = merge_completion_params(llm_cfg.get("completion_params") or {}, request)
if not isinstance(params, dict):
params = {}
llm_cfg["completion_params"] = params
# Host ModelInvokeLLMRequest requires model_parameters: dict.
# Keep both keys so either mapping path works.
llm_cfg["model_parameters"] = params
return llm_cfg

def merge_tool_calls(existing: Any, incoming: Any) -> list[Any]:
"""Merge tool_calls by id (preferred) or index; last non-empty fields win."""
if not incoming:
return list(existing or [])
if not existing:
return list(incoming)

merged: list[Any] = list(existing)
by_id: dict[str, int] = {}
for i, tc in enumerate(merged):
tc_id = getattr(tc, "id", None)
if tc_id is None and isinstance(tc, dict):
tc_id = tc.get("id")
if tc_id:
by_id[str(tc_id)] = i

for index, tc in enumerate(incoming):
tc_id = getattr(tc, "id", None)
if tc_id is None and isinstance(tc, dict):
tc_id = tc.get("id")
if tc_id and str(tc_id) in by_id:
merged[by_id[str(tc_id)]] = tc
continue
if index < len(merged) and not tc_id:
# index-aligned partial update without id
merged[index] = tc
continue
merged.append(tc)
if tc_id:
by_id[str(tc_id)] = len(merged) - 1
return merged

def tool_calls_to_openai(tool_calls: Any) -> list[dict[str, Any]]:
if not tool_calls:
return []
result: list[dict[str, Any]] = []
for index, tool_call in enumerate(tool_calls):
function = getattr(tool_call, "function", None)
name = getattr(function, "name", "") if function is not None else ""
arguments = getattr(function, "arguments", "") if function is not None else ""
result.append(
{
"index": index,
"id": getattr(tool_call, "id", "") or f"call_{index}",
"type": getattr(tool_call, "type", None) or "function",
"function": {
"name": name or "",
"arguments": _normalize_arguments(arguments),
},
}
)
return result

def usage_to_openai(usage: Any) -> dict[str, int]:
if usage is None:
return {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
return {
"prompt_tokens": int(getattr(usage, "prompt_tokens", 0) or 0),
"completion_tokens": int(getattr(usage, "completion_tokens", 0) or 0),
"total_tokens": int(getattr(usage, "total_tokens", 0) or 0),
}

def infer_finish_reason(tool_calls: Any) -> str:
return "tool_calls" if tool_calls else "stop"

def new_completion_id() -> str:
return "chatcmpl-" + str(uuid.uuid4())

def build_chat_completion_response(
*,
model: str | None,
message_content: Any,
tool_calls: Any,
usage: Any,
completion_id: str,
created: int | None = None,
) -> dict[str, Any]:
openai_tool_calls = tool_calls_to_openai(tool_calls)
finish_reason = infer_finish_reason(openai_tool_calls)
# Normalize to OpenAI text (or null). Never Python-repr objects/lists.
if isinstance(message_content, str):
content: Any = message_content
elif message_content is None:
content = None
else:
content = _normalize_stream_content(message_content)
if openai_tool_calls and (content is None or content == ""):
content = None
message: dict[str, Any] = {
"role": "assistant",
"content": content,
}
if openai_tool_calls:
message["tool_calls"] = [
{
"id": item["id"],
"type": item["type"],
"function": item["function"],
}
for item in openai_tool_calls
]
return {
"id": completion_id,
"object": "chat.completion",
"created": int(created if created is not None else time.time()),
"model": model,
"system_fingerprint": f"oaicompat-{PLUGIN_VERSION}-{PLUGIN_RUNTIME_FINGERPRINT}",
"choices": [
{
"index": 0,
"message": message,
"finish_reason": finish_reason,
}
],
"usage": usage_to_openai(usage),
}

def _sse_line(payload: dict[str, Any]) -> str:
return f"data: {json.dumps(payload, ensure_ascii=False)}\n\n"

def _chunk_payload(
*,
completion_id: str,
created: int,
model: str | None,
delta: dict[str, Any],
finish_reason: str | None = None,
usage: dict[str, int] | None = None,
) -> dict[str, Any]:
body: dict[str, Any] = {
"id": completion_id,
"object": "chat.completion.chunk",
"created": created,
"model": model,
"system_fingerprint": f"oaicompat-{PLUGIN_VERSION}-{PLUGIN_RUNTIME_FINGERPRINT}",
"choices": [
{
"index": 0,
"delta": delta,
"finish_reason": finish_reason,
}
],
}
if usage is not None:
body["usage"] = usage
return body

def _normalize_stream_content(content: Any) -> str | None:
if content is None:
return None
if isinstance(content, str):
return content or None
if isinstance(content, list):
parts: list[str] = []
for item in content:
if isinstance(item, dict):
item_type = item.get("type")
if item_type in (None, "text"):
text = item.get("text")
if text is None:
text = item.get("data")
if text is not None:
parts.append(str(text))
continue
data = getattr(item, "data", None)
if data is not None:
parts.append(str(data))
continue
text = getattr(item, "text", None)
if text is not None:
parts.append(str(text))
joined = "".join(parts)
return joined or None
return str(content) or None

def iter_sse_chat_completion(
*,
model: str | None,
chunk_iter: Iterable[Any],
completion_id: str,
created: int | None = None,
) -> Iterator[str]:
"""Emit OpenAI-compatible SSE frames.

Always ends with a finish_reason frame + `data: [DONE]`, even if the
upstream chunk iterator raises. Clients like Codex abort when a stream
ends without finish_reason.
"""
created_ts = int(created if created is not None else time.time())
yield _sse_line(
_chunk_payload(
completion_id=completion_id,
created=created_ts,
model=model,
delta={"role": "assistant", "content": ""},
finish_reason=None,
)
)

last_tool_calls: Any = None
last_finish: str | None = None
last_usage: Any = None
stream_error: Exception | None = None
finish = "stop"
openai_tool_calls: list[dict[str, Any]] = []

try:
for chunk in chunk_iter:
delta = getattr(chunk, "delta", None)
if delta is None:
continue
message = getattr(delta, "message", None)
content = (
getattr(message, "content", None) if message is not None else None
)
text = _normalize_stream_content(content)
if text:
yield _sse_line(
_chunk_payload(
completion_id=completion_id,
created=created_ts,
model=model,
delta={"content": text},
finish_reason=None,
)
)
tool_calls = (
getattr(message, "tool_calls", None) if message is not None else None
)
if tool_calls:
last_tool_calls = merge_tool_calls(last_tool_calls, tool_calls)
finish_reason = getattr(delta, "finish_reason", None)
if finish_reason:
last_finish = finish_reason
usage = getattr(delta, "usage", None)
if usage is not None:
last_usage = usage

try:
openai_tool_calls = tool_calls_to_openai(last_tool_calls)
except Exception as exc: # noqa: BLE001
stream_error = exc
openai_tool_calls = []

if openai_tool_calls:
yield _sse_line(
_chunk_payload(
completion_id=completion_id,
created=created_ts,
model=model,
delta={"tool_calls": openai_tool_calls},
finish_reason=None,
)
)
finish = "tool_calls"
elif last_finish in {"stop", "length", "content_filter", "tool_calls"}:
finish = last_finish
else:
finish = "stop"
except Exception as exc: # noqa: BLE001 - must still close SSE for clients
stream_error = exc
# Prefer tool_calls finish if we already observed tools before the error.
finish = "tool_calls" if last_tool_calls else "stop"

# If we collected tool_calls but the success path did not emit them (exception
# mid-loop), emit one final tool_calls frame before finish.
if last_tool_calls and not openai_tool_calls:
try:
openai_tool_calls = tool_calls_to_openai(last_tool_calls)
except Exception as exc: # noqa: BLE001
stream_error = stream_error or exc
openai_tool_calls = []
if openai_tool_calls:
yield _sse_line(
_chunk_payload(
completion_id=completion_id,
created=created_ts,
model=model,
delta={"tool_calls": openai_tool_calls},
finish_reason=None,
)
)
finish = "tool_calls"

finish_payload = _chunk_payload(
completion_id=completion_id,
created=created_ts,
model=model,
delta={},
finish_reason=finish,
usage=usage_to_openai(last_usage),
)
if stream_error is not None:
# Keep OpenAI shape; surface upstream failure without dropping finish_reason.
finish_payload["error"] = openai_error(
str(stream_error) or "upstream stream error",
type_="server_error",
)["error"]
yield _sse_line(finish_payload)
yield "data: [DONE]\n\n"

def _role_value(role: Any) -> str:
if role is None:
return ""
value = getattr(role, "value", role)
if not isinstance(value, str):
value = str(value)
return value.strip().lower()

def serialize_prompt_message(message: PromptMessage) -> dict[str, Any]:
"""Serialize PromptMessage for host ModelInvokeLLMRequest.

Always emit plain JSON-safe values. Never leave Enum objects or empty roles.
"""
role = _role_value(getattr(message, "role", None))
if role not in {"system", "user", "assistant", "tool", "developer"}:
raise ProtocolError(
f"invalid message role after parse: {getattr(message, 'role', None)!r}"
)

# Dify host (Go daemon) rejects null content; use "" so assistant messages
# with tool_calls still parse. See sanitize_host_llm_payload.
content = getattr(message, "content", None)
if content is None:
content = ""
elif not isinstance(content, (str, list)):
content = str(content)

payload: dict[str, Any] = {
"role": role,
"content": content,
}
name = getattr(message, "name", None)
if name:
payload["name"] = str(name)

if role == "assistant":
# Host RequestInvokeLLM requires tool_calls to be a list, never null.
raw_tool_calls = getattr(message, "tool_calls", None) or []
serialized_calls: list[dict[str, Any]] = []
for index, tool_call in enumerate(raw_tool_calls):
function = getattr(tool_call, "function", None)
serialized_calls.append(
{
"id": str(getattr(tool_call, "id", "") or f"call_{index}"),
"type": str(getattr(tool_call, "type", None) or "function"),
"function": {
"name": str(getattr(function, "name", "") or ""),
"arguments": _normalize_arguments(
getattr(function, "arguments", "") if function else ""
),
},
}
)
payload["tool_calls"] = serialized_calls
elif role == "tool":
tool_call_id = getattr(message, "tool_call_id", None)
if not tool_call_id:
raise ProtocolError("tool message requires tool_call_id")
payload["tool_call_id"] = str(tool_call_id)

return payload

def serialize_tool(tool: PromptMessageTool | dict[str, Any]) -> dict[str, Any]:
if isinstance(tool, dict):
name = tool.get("name")
if not name:
raise ProtocolError("tool.name is required")
return {
"name": str(name),
"description": str(tool.get("description") or ""),
"parameters": tool.get("parameters") or {},
}
return {
"name": str(tool.name),
"description": str(tool.description or ""),
"parameters": tool.parameters or {},
}

def build_host_llm_invoke_payload(
*,
llm_cfg: dict[str, Any],
prompt_messages: list[PromptMessage],
tools: list[PromptMessageTool] | None,
stop: list[str] | None,
stream: bool,
) -> dict[str, Any]:
"""Build reverse-invoke payload that passes host ModelInvokeLLMRequest.

Why this exists:
- SDK LLMInvocation does `**model_config.model_dump()` + message.model_dump().
- Host validates `model_parameters: dict` (not completion_params).
- Host rejects empty role strings.
- Enum objects in nested dumps can serialize poorly depending on path.

This builder emits only JSON-safe primitives with dual-written params.
"""
if not isinstance(llm_cfg, dict):
raise ProtocolError("LLM config is invalid")

provider = llm_cfg.get("provider")
model = llm_cfg.get("model")
mode = llm_cfg.get("mode") or "chat"
if not provider or not model:
raise ProtocolError("LLM provider/model is required")

params = llm_cfg.get("model_parameters")
if not isinstance(params, dict):
params = llm_cfg.get("completion_params")
if not isinstance(params, dict):
params = {}

model_type = llm_cfg.get("model_type") or "llm"
if hasattr(model_type, "value"):
model_type = model_type.value
model_type = str(model_type)

serialized_messages = [serialize_prompt_message(m) for m in prompt_messages]
if not serialized_messages:
raise ProtocolError("messages must be a non-empty list")
for item in serialized_messages:
if not item.get("role"):
raise ProtocolError("serialized message role is empty")

serialized_tools = None
if tools:
serialized_tools = [serialize_tool(t) for t in tools]

# Dual-write both keys. Host currently requires model_parameters.
payload = {
"provider": str(provider),
"model": str(model),
"model_type": model_type,
"mode": str(mode),
"completion_params": params,
"model_parameters": params,
"prompt_messages": serialized_messages,
"tools": serialized_tools,
"stop": stop,
"stream": bool(stream),
}
return sanitize_host_llm_payload(payload)

def sanitize_host_llm_payload(payload: dict[str, Any]) -> dict[str, Any]:
"""Final hard sanitization before reverse-invoke.

Guarantees host ModelInvokeLLMRequest / RequestInvokeLLM constraints:
- model_parameters is a dict (never None)
- completion_params is a dict
- every prompt_messages[].role is a non-empty allowed string
- every assistant message has tool_calls as a list (never null)
- tools is list or None (not other types)
- stop is list[str] or None
"""
if not isinstance(payload, dict):
raise ProtocolError("invoke payload must be a dict")

out = dict(payload)

params = out.get("model_parameters")
if not isinstance(params, dict):
params = out.get("completion_params")
if not isinstance(params, dict):
params = {}
out["model_parameters"] = params
out["completion_params"] = params

provider = out.get("provider")
model = out.get("model")
if not provider or not model:
raise ProtocolError("LLM provider/model is required")
out["provider"] = str(provider)
out["model"] = str(model)
mode = out.get("mode") or "chat"
out["mode"] = str(mode)
model_type = out.get("model_type") or "llm"
if hasattr(model_type, "value"):
model_type = model_type.value
out["model_type"] = str(model_type)

messages = out.get("prompt_messages")
if not isinstance(messages, list) or not messages:
raise ProtocolError("messages must be a non-empty list")

allowed = {"system", "user", "assistant", "tool", "developer"}
cleaned_messages: list[dict[str, Any]] = []
for item in messages:
if not isinstance(item, dict):
raise ProtocolError("each prompt message must be an object")
role_raw = item.get("role")
if hasattr(role_raw, "value"):
role_raw = role_raw.value
if not isinstance(role_raw, str) or not role_raw.strip():
raise ProtocolError(
"message role is required and must be one of: system, user, assistant, tool"
)
role = role_raw.strip().lower()
if role not in allowed:
raise ProtocolError(f"invalid message role: {role_raw}")

# Dify host (Go daemon) hard-rejects null content: its PromptMessage
# JSON decoder returns "content field is required" for null, which fails
# the whole reverse-invoke parse and degrades every field (role='',
# model_parameters=None) at the model provider. Native Dify uses "".
raw_content = item.get("content")
if raw_content is None:
raw_content = ""
elif not isinstance(raw_content, (str, list)):
raw_content = str(raw_content)
msg: dict[str, Any] = {
"role": role,
"content": raw_content,
}
name = item.get("name")
if name:
msg["name"] = str(name)

if role == "assistant":
tool_calls = item.get("tool_calls")
if tool_calls is None:
tool_calls = []
if not isinstance(tool_calls, list):
raise ProtocolError("assistant tool_calls must be a list")
fixed_calls: list[dict[str, Any]] = []
for index, call in enumerate(tool_calls):
if not isinstance(call, dict):
raise ProtocolError("tool_call must be an object")
function = call.get("function") or {}
if not isinstance(function, dict):
raise ProtocolError("tool_call.function must be an object")
fixed_calls.append(
{
"id": str(call.get("id") or f"call_{index}"),
"type": str(call.get("type") or "function"),
"function": {
"name": str(function.get("name") or ""),
"arguments": _normalize_arguments(function.get("arguments")),
},
}
)
msg["tool_calls"] = fixed_calls
elif role == "tool":
tool_call_id = item.get("tool_call_id")
if not tool_call_id:
raise ProtocolError("tool message requires tool_call_id")
msg["tool_call_id"] = str(tool_call_id)

cleaned_messages.append(msg)

out["prompt_messages"] = cleaned_messages

tools = out.get("tools")
if tools is None:
out["tools"] = None
elif isinstance(tools, list):
cleaned_tools: list[dict[str, Any]] = []
for tool in tools:
if not isinstance(tool, dict):
raise ProtocolError("each tool must be an object")
name = tool.get("name")
if not name:
raise ProtocolError("tool.name is required")
cleaned_tools.append(
{
"name": str(name),
"description": str(tool.get("description") or ""),
"parameters": tool.get("parameters")
if isinstance(tool.get("parameters"), dict)
else {},
}
)
out["tools"] = cleaned_tools
else:
raise ProtocolError("tools must be a list or null")

stop = out.get("stop")
if stop is None:
out["stop"] = None
elif isinstance(stop, list) and all(isinstance(s, str) for s in stop):
out["stop"] = list(stop)
else:
raise ProtocolError("stop must be string list or null")

out["stream"] = bool(out.get("stream", False))
return out



extensions/oaicompat_dify_model/endpoints/text_embedding.py



点击展开
import json
from collections.abc import Mapping
from typing import Optional

from dify_plugin import Endpoint
from dify_plugin.entities.model.text_embedding import TextEmbeddingModelConfig
from werkzeug import Request, Response

from endpoints.auth import BaseAuth
from endpoints.openai_protocol import ProtocolError, openai_error

class OaicompatDifyModelEndpoint(Endpoint, BaseAuth):
def _error(self, status: int, message: str, type_: str = "invalid_request_error") -> Response:
return Response(
json.dumps(openai_error(message, type_=type_)),
status=status,
content_type="application/json",
)

def _invoke(self, r: Request, values: Mapping, settings: Mapping) -> Response:
if not self.verify(r, settings):
return self._error(401, "Unauthorized")

try:
model: Optional[dict] = settings.get("text_embedding", None)
if not model:
raise ProtocolError("Text embedding model is not set")

try:
data = r.get_json(force=True)
except Exception as exc: # noqa: BLE001
raise ProtocolError(f"invalid JSON body: {exc}") from exc
if not data:
raise ProtocolError("Request body is empty")

texts: list[str] = []
if isinstance(data.get("input"), str):
texts.append(data.get("input"))
elif isinstance(data.get("input"), list):
texts = data.get("input")
else:
raise ProtocolError("Invalid input type")

text_embedding_response = self.session.model.text_embedding.invoke(
model_config=TextEmbeddingModelConfig(**model),
texts=texts,
)

return Response(
json.dumps(
{
"object": "list",
"data": [
{
"object": "embedding",
"embedding": embedding,
"index": index,
}
for index, embedding in enumerate(
text_embedding_response.embeddings
)
],
"usage": {
"prompt_tokens": text_embedding_response.usage.total_tokens,
"total_tokens": text_embedding_response.usage.total_tokens,
},
"model": text_embedding_response.model,
}
),
status=200,
content_type="application/json",
)
except ProtocolError as exc:
return self._error(400, str(exc))
except ValueError as exc:
return self._error(400, str(exc))
except Exception as exc: # noqa: BLE001
return self._error(500, str(exc) or type(exc).__name__, "server_error")



至于打包成dify插件,需要安装dify相关包(自行研究下哈)


不想折腾的佬,可以直接使用已经打包好的插件哈


oaicompat_dify_model.difypkg.zip (100.2 KB)

最新回复 (5)
  • 山海 08-15 13:19
    1

    这个大概有多少额度呀,我记得我之前也有一个,一直没用过

  • ctrol9 楼主 08-15 13:29
    2

    每个月5000积分,用起来其实一个号是不怎么够用的…

  • jdjdk 08-15 13:44
    4

    请问dify的模型纯血吗,还是加了亿些系统提示的那种呀

  • ctrol9 楼主 08-15 13:50
    5

    是不是纯血不确定,但是目前用起来没有发现有额外提示词的情况,体感也跟官方的差不多(点名kiro ^-^)

  • jdjdk 08-15 13:52
    6

    okk,感谢大佬,要是搓出就好了哈哈哈cc兼容的anthropic格式就完美了,我研究一下下

* 帖子来源Linux.do
返回