mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
fix: improve Service API OpenAPI contracts (#37592)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
autofix-ci[bot]
parent
9021b3f5be
commit
baf775134e
@@ -1,7 +1,8 @@
|
||||
from typing import Any, Literal
|
||||
from copy import deepcopy
|
||||
from typing import Any, Literal, override
|
||||
from uuid import UUID
|
||||
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
from pydantic import BaseModel, Field, GetJsonSchemaHandler, model_validator
|
||||
|
||||
from libs.helper import UUIDStrOrEmpty
|
||||
|
||||
@@ -12,6 +13,45 @@ class ConversationRenamePayload(BaseModel):
|
||||
name: str | None = None
|
||||
auto_generate: bool = False
|
||||
|
||||
@classmethod
|
||||
@override
|
||||
def __get_pydantic_json_schema__(cls, core_schema: Any, handler: GetJsonSchemaHandler) -> dict[str, Any]:
|
||||
schema = handler.resolve_ref_schema(handler(core_schema))
|
||||
properties = schema.get("properties")
|
||||
if not isinstance(properties, dict):
|
||||
return schema
|
||||
|
||||
auto_generate_schema = deepcopy(properties.get("auto_generate", {"type": "boolean"}))
|
||||
name_schema = deepcopy(properties.get("name", {"type": "string"}))
|
||||
non_blank_name_schema: dict[str, Any] = {"pattern": r".*\S.*", "type": "string"}
|
||||
if isinstance(name_schema, dict) and isinstance(name_schema.get("title"), str):
|
||||
non_blank_name_schema["title"] = name_schema["title"]
|
||||
|
||||
auto_generate_true_schema = {**auto_generate_schema, "enum": [True]}
|
||||
auto_generate_true_schema.pop("default", None)
|
||||
|
||||
return {
|
||||
**schema,
|
||||
"anyOf": [
|
||||
{
|
||||
"properties": {
|
||||
"auto_generate": auto_generate_true_schema,
|
||||
"name": name_schema,
|
||||
},
|
||||
"required": ["auto_generate"],
|
||||
"type": "object",
|
||||
},
|
||||
{
|
||||
"properties": {
|
||||
"auto_generate": {**auto_generate_schema, "enum": [False]},
|
||||
"name": non_blank_name_schema,
|
||||
},
|
||||
"required": ["name"],
|
||||
"type": "object",
|
||||
},
|
||||
],
|
||||
}
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_name_requirement(self):
|
||||
if not self.auto_generate:
|
||||
|
||||
@@ -45,6 +45,13 @@ class AnnotationJobStatusResponse(ResponseModel):
|
||||
error_msg: str | None = None
|
||||
|
||||
|
||||
ANNOTATION_REPLY_ACTION_PARAM = {
|
||||
"description": "Action to perform: 'enable' or 'disable'",
|
||||
"enum": ["enable", "disable"],
|
||||
"type": "string",
|
||||
}
|
||||
|
||||
|
||||
register_schema_models(
|
||||
service_api_ns,
|
||||
AnnotationCreatePayload,
|
||||
@@ -61,7 +68,7 @@ class AnnotationReplyActionApi(Resource):
|
||||
@service_api_ns.expect(service_api_ns.models[AnnotationReplyActionPayload.__name__])
|
||||
@service_api_ns.doc("annotation_reply_action")
|
||||
@service_api_ns.doc(description="Enable or disable annotation reply feature")
|
||||
@service_api_ns.doc(params={"action": "Action to perform: 'enable' or 'disable'"})
|
||||
@service_api_ns.doc(params={"action": ANNOTATION_REPLY_ACTION_PARAM})
|
||||
@service_api_ns.doc(
|
||||
responses={
|
||||
200: "Action completed successfully",
|
||||
|
||||
@@ -20,6 +20,7 @@ from controllers.service_api.app.error import (
|
||||
ProviderQuotaExceededError,
|
||||
UnsupportedAudioTypeError,
|
||||
)
|
||||
from controllers.service_api.schema import binary_response, expect_with_user, multipart_file_params
|
||||
from controllers.service_api.wraps import FetchUserArg, WhereisUserArg, validate_app_token
|
||||
from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError
|
||||
from graphon.model_runtime.errors.invoke import InvokeError
|
||||
@@ -41,6 +42,7 @@ register_response_schema_models(service_api_ns, AudioBinaryResponse, AudioTransc
|
||||
class AudioApi(Resource):
|
||||
@service_api_ns.doc("audio_to_text")
|
||||
@service_api_ns.doc(description="Convert audio to text using speech-to-text")
|
||||
@service_api_ns.doc(consumes=["multipart/form-data"], params=multipart_file_params(include_user=True))
|
||||
@service_api_ns.doc(
|
||||
responses={
|
||||
200: "Audio successfully transcribed",
|
||||
@@ -99,7 +101,8 @@ register_schema_model(service_api_ns, TextToAudioPayload)
|
||||
|
||||
@service_api_ns.route("/text-to-audio")
|
||||
class TextApi(Resource):
|
||||
@service_api_ns.expect(service_api_ns.models[TextToAudioPayload.__name__])
|
||||
@expect_with_user(service_api_ns, TextToAudioPayload)
|
||||
@binary_response(service_api_ns, "audio/mpeg")
|
||||
@service_api_ns.doc("text_to_audio")
|
||||
@service_api_ns.doc(description="Convert text to audio using text-to-speech")
|
||||
@service_api_ns.doc(
|
||||
@@ -110,11 +113,7 @@ class TextApi(Resource):
|
||||
500: "Internal server error",
|
||||
}
|
||||
)
|
||||
@service_api_ns.response(
|
||||
200,
|
||||
"Text successfully converted to audio",
|
||||
service_api_ns.models[AudioBinaryResponse.__name__],
|
||||
)
|
||||
@service_api_ns.response(200, "Text successfully converted to audio")
|
||||
@validate_app_token(fetch_user_arg=FetchUserArg(fetch_from=WhereisUserArg.JSON))
|
||||
def post(self, app_model: App, end_user: EndUser):
|
||||
"""Convert text to audio using text-to-speech.
|
||||
|
||||
@@ -20,6 +20,7 @@ from controllers.service_api.app.error import (
|
||||
ProviderNotInitializeError,
|
||||
ProviderQuotaExceededError,
|
||||
)
|
||||
from controllers.service_api.schema import expect_user_json, expect_with_user, json_or_event_stream_response
|
||||
from controllers.service_api.wraps import FetchUserArg, WhereisUserArg, validate_app_token
|
||||
from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||
@@ -92,7 +93,8 @@ register_response_schema_models(service_api_ns, GeneratedAppResponse, SimpleResu
|
||||
|
||||
@service_api_ns.route("/completion-messages")
|
||||
class CompletionApi(Resource):
|
||||
@service_api_ns.expect(service_api_ns.models[CompletionRequestPayload.__name__])
|
||||
@expect_with_user(service_api_ns, CompletionRequestPayload)
|
||||
@json_or_event_stream_response(service_api_ns)
|
||||
@service_api_ns.doc("create_completion")
|
||||
@service_api_ns.doc(description="Create a completion for the given prompt")
|
||||
@service_api_ns.doc(
|
||||
@@ -168,6 +170,7 @@ class CompletionApi(Resource):
|
||||
|
||||
@service_api_ns.route("/completion-messages/<string:task_id>/stop")
|
||||
class CompletionStopApi(Resource):
|
||||
@expect_user_json(service_api_ns)
|
||||
@service_api_ns.doc("stop_completion")
|
||||
@service_api_ns.doc(description="Stop a running completion task")
|
||||
@service_api_ns.doc(params={"task_id": "The ID of the task to stop"})
|
||||
@@ -197,7 +200,8 @@ class CompletionStopApi(Resource):
|
||||
|
||||
@service_api_ns.route("/chat-messages")
|
||||
class ChatApi(Resource):
|
||||
@service_api_ns.expect(service_api_ns.models[ChatRequestPayload.__name__])
|
||||
@expect_with_user(service_api_ns, ChatRequestPayload)
|
||||
@json_or_event_stream_response(service_api_ns)
|
||||
@service_api_ns.doc("create_chat_message")
|
||||
@service_api_ns.doc(description="Send a message in a chat conversation")
|
||||
@service_api_ns.doc(
|
||||
@@ -276,6 +280,7 @@ class ChatApi(Resource):
|
||||
|
||||
@service_api_ns.route("/chat-messages/<string:task_id>/stop")
|
||||
class ChatStopApi(Resource):
|
||||
@expect_user_json(service_api_ns)
|
||||
@service_api_ns.doc("stop_chat_message")
|
||||
@service_api_ns.doc(description="Stop a running chat message generation")
|
||||
@service_api_ns.doc(params={"task_id": "The ID of the task to stop"})
|
||||
|
||||
@@ -13,6 +13,7 @@ from controllers.common.controller_schemas import ConversationRenamePayload
|
||||
from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models
|
||||
from controllers.service_api import service_api_ns
|
||||
from controllers.service_api.app.error import NotChatAppError
|
||||
from controllers.service_api.schema import expect_user_json, expect_with_user
|
||||
from controllers.service_api.wraps import FetchUserArg, WhereisUserArg, validate_app_token
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||
from extensions.ext_database import db
|
||||
@@ -197,6 +198,7 @@ class ConversationApi(Resource):
|
||||
|
||||
@service_api_ns.route("/conversations/<uuid:c_id>")
|
||||
class ConversationDetailApi(Resource):
|
||||
@expect_user_json(service_api_ns)
|
||||
@service_api_ns.doc("delete_conversation")
|
||||
@service_api_ns.doc(description="Delete a specific conversation")
|
||||
@service_api_ns.doc(params={"c_id": "Conversation ID"})
|
||||
@@ -225,7 +227,7 @@ class ConversationDetailApi(Resource):
|
||||
|
||||
@service_api_ns.route("/conversations/<uuid:c_id>/name")
|
||||
class ConversationRenameApi(Resource):
|
||||
@service_api_ns.expect(service_api_ns.models[ConversationRenamePayload.__name__])
|
||||
@expect_with_user(service_api_ns, ConversationRenamePayload)
|
||||
@service_api_ns.doc("rename_conversation")
|
||||
@service_api_ns.doc(description="Rename a conversation or auto-generate a name")
|
||||
@service_api_ns.doc(params={"c_id": "Conversation ID"})
|
||||
@@ -312,7 +314,7 @@ class ConversationVariablesApi(Resource):
|
||||
|
||||
@service_api_ns.route("/conversations/<uuid:c_id>/variables/<uuid:variable_id>")
|
||||
class ConversationVariableDetailApi(Resource):
|
||||
@service_api_ns.expect(service_api_ns.models[ConversationVariableUpdatePayload.__name__])
|
||||
@expect_with_user(service_api_ns, ConversationVariableUpdatePayload)
|
||||
@service_api_ns.doc("update_conversation_variable")
|
||||
@service_api_ns.doc(description="Update a conversation variable's value")
|
||||
@service_api_ns.doc(params={"c_id": "Conversation ID", "variable_id": "Variable ID"})
|
||||
|
||||
@@ -12,6 +12,7 @@ from controllers.common.errors import (
|
||||
)
|
||||
from controllers.common.schema import register_schema_models
|
||||
from controllers.service_api import service_api_ns
|
||||
from controllers.service_api.schema import multipart_file_params
|
||||
from controllers.service_api.wraps import FetchUserArg, WhereisUserArg, validate_app_token
|
||||
from extensions.ext_database import db
|
||||
from fields.file_fields import FileResponse
|
||||
@@ -25,6 +26,7 @@ register_schema_models(service_api_ns, FileResponse)
|
||||
class FileApi(Resource):
|
||||
@service_api_ns.doc("upload_file")
|
||||
@service_api_ns.doc(description="Upload a file for use in conversations")
|
||||
@service_api_ns.doc(consumes=["multipart/form-data"], params=multipart_file_params(include_user=True))
|
||||
@service_api_ns.doc(
|
||||
responses={
|
||||
201: "File uploaded successfully",
|
||||
|
||||
@@ -15,6 +15,7 @@ from controllers.service_api.app.error import (
|
||||
FileAccessDeniedError,
|
||||
FileNotFoundError,
|
||||
)
|
||||
from controllers.service_api.schema import binary_response
|
||||
from controllers.service_api.wraps import FetchUserArg, WhereisUserArg, validate_app_token
|
||||
from extensions.ext_database import db
|
||||
from extensions.ext_storage import storage
|
||||
@@ -30,6 +31,26 @@ class FilePreviewQuery(BaseModel):
|
||||
register_schema_model(service_api_ns, FilePreviewQuery)
|
||||
register_response_schema_model(service_api_ns, BinaryFileResponse)
|
||||
|
||||
FILE_PREVIEW_RESPONSE_MEDIA_TYPES = [
|
||||
"application/octet-stream",
|
||||
"application/pdf",
|
||||
"audio/aac",
|
||||
"audio/flac",
|
||||
"audio/mp4",
|
||||
"audio/mpeg",
|
||||
"audio/ogg",
|
||||
"audio/wav",
|
||||
"audio/x-m4a",
|
||||
"image/gif",
|
||||
"image/jpeg",
|
||||
"image/png",
|
||||
"image/webp",
|
||||
"text/plain",
|
||||
"video/mp4",
|
||||
"video/quicktime",
|
||||
"video/webm",
|
||||
]
|
||||
|
||||
|
||||
@service_api_ns.route("/files/<uuid:file_id>/preview")
|
||||
class FilePreviewApi(Resource):
|
||||
@@ -41,6 +62,7 @@ class FilePreviewApi(Resource):
|
||||
"""
|
||||
|
||||
@service_api_ns.doc(params=query_params_from_model(FilePreviewQuery))
|
||||
@binary_response(service_api_ns, FILE_PREVIEW_RESPONSE_MEDIA_TYPES)
|
||||
@service_api_ns.doc("preview_file")
|
||||
@service_api_ns.doc(description="Preview or download a file uploaded via Service API")
|
||||
@service_api_ns.doc(params={"file_id": "UUID of the file to preview"})
|
||||
@@ -52,11 +74,7 @@ class FilePreviewApi(Resource):
|
||||
404: "File not found",
|
||||
}
|
||||
)
|
||||
@service_api_ns.response(
|
||||
200,
|
||||
"File retrieved successfully",
|
||||
service_api_ns.models[BinaryFileResponse.__name__],
|
||||
)
|
||||
@service_api_ns.response(200, "File retrieved successfully")
|
||||
@validate_app_token(fetch_user_arg=FetchUserArg(fetch_from=WhereisUserArg.QUERY))
|
||||
def get(self, app_model: App, end_user: EndUser, file_id: UUID):
|
||||
"""
|
||||
|
||||
@@ -18,6 +18,7 @@ from werkzeug.exceptions import BadRequest, NotFound
|
||||
from controllers.common.human_input import HumanInputFormSubmitPayload, stringify_form_default_values
|
||||
from controllers.common.schema import register_response_schema_models, register_schema_models
|
||||
from controllers.service_api import service_api_ns
|
||||
from controllers.service_api.schema import expect_with_user
|
||||
from controllers.service_api.wraps import FetchUserArg, WhereisUserArg, validate_app_token
|
||||
from core.workflow.human_input_policy import HumanInputSurface, is_recipient_type_allowed_for_surface
|
||||
from extensions.ext_database import db
|
||||
@@ -101,7 +102,7 @@ class WorkflowHumanInputFormApi(Resource):
|
||||
inputs = service.resolve_form_inputs(form)
|
||||
return _jsonify_form_definition(form, inputs=inputs)
|
||||
|
||||
@service_api_ns.expect(service_api_ns.models[HumanInputFormSubmitPayload.__name__])
|
||||
@expect_with_user(service_api_ns, HumanInputFormSubmitPayload)
|
||||
@service_api_ns.doc("submit_human_input_form")
|
||||
@service_api_ns.doc(description="Submit a paused human input form by token")
|
||||
@service_api_ns.doc(params={"form_token": "Human input form token"})
|
||||
|
||||
@@ -12,6 +12,7 @@ from controllers.common.fields import SimpleResultStringListResponse
|
||||
from controllers.common.schema import query_params_from_model, register_response_schema_models, register_schema_models
|
||||
from controllers.service_api import service_api_ns
|
||||
from controllers.service_api.app.error import NotChatAppError
|
||||
from controllers.service_api.schema import expect_with_user
|
||||
from controllers.service_api.wraps import FetchUserArg, WhereisUserArg, validate_app_token
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||
from fields.base import ResponseModel
|
||||
@@ -112,7 +113,7 @@ class MessageListApi(Resource):
|
||||
|
||||
@service_api_ns.route("/messages/<uuid:message_id>/feedbacks")
|
||||
class MessageFeedbackApi(Resource):
|
||||
@service_api_ns.expect(service_api_ns.models[MessageFeedbackPayload.__name__])
|
||||
@expect_with_user(service_api_ns, MessageFeedbackPayload)
|
||||
@service_api_ns.response(200, "Feedback submitted successfully", service_api_ns.models[ResultResponse.__name__])
|
||||
@service_api_ns.doc("create_message_feedback")
|
||||
@service_api_ns.doc(description="Submit feedback for a message")
|
||||
|
||||
@@ -21,6 +21,11 @@ from controllers.service_api.app.error import (
|
||||
ProviderNotInitializeError,
|
||||
ProviderQuotaExceededError,
|
||||
)
|
||||
from controllers.service_api.schema import (
|
||||
expect_user_json,
|
||||
expect_with_user,
|
||||
json_or_event_stream_response,
|
||||
)
|
||||
from controllers.service_api.wraps import FetchUserArg, WhereisUserArg, validate_app_token
|
||||
from controllers.web.error import InvokeRateLimitError as InvokeRateLimitHttpError
|
||||
from core.app.apps.base_app_queue_manager import AppQueueManager
|
||||
@@ -249,7 +254,8 @@ class WorkflowRunDetailApi(Resource):
|
||||
|
||||
@service_api_ns.route("/workflows/run")
|
||||
class WorkflowRunApi(Resource):
|
||||
@service_api_ns.expect(service_api_ns.models[WorkflowRunPayload.__name__])
|
||||
@expect_with_user(service_api_ns, WorkflowRunPayload)
|
||||
@json_or_event_stream_response(service_api_ns)
|
||||
@service_api_ns.doc("run_workflow")
|
||||
@service_api_ns.doc(description="Execute a workflow")
|
||||
@service_api_ns.doc(
|
||||
@@ -313,7 +319,8 @@ class WorkflowRunApi(Resource):
|
||||
|
||||
@service_api_ns.route("/workflows/<string:workflow_id>/run")
|
||||
class WorkflowRunByIdApi(Resource):
|
||||
@service_api_ns.expect(service_api_ns.models[WorkflowRunPayload.__name__])
|
||||
@expect_with_user(service_api_ns, WorkflowRunPayload)
|
||||
@json_or_event_stream_response(service_api_ns)
|
||||
@service_api_ns.doc("run_workflow_by_id")
|
||||
@service_api_ns.doc(description="Execute a specific workflow by ID")
|
||||
@service_api_ns.doc(params={"workflow_id": "Workflow ID to execute"})
|
||||
@@ -387,6 +394,7 @@ class WorkflowRunByIdApi(Resource):
|
||||
|
||||
@service_api_ns.route("/workflows/tasks/<string:task_id>/stop")
|
||||
class WorkflowTaskStopApi(Resource):
|
||||
@expect_user_json(service_api_ns)
|
||||
@service_api_ns.doc("stop_workflow_task")
|
||||
@service_api_ns.doc(description="Stop a running workflow task")
|
||||
@service_api_ns.doc(params={"task_id": "Task ID to stop"})
|
||||
|
||||
@@ -15,6 +15,7 @@ from controllers.common.fields import EventStreamResponse
|
||||
from controllers.common.schema import query_params_from_model, register_response_schema_model, register_schema_models
|
||||
from controllers.service_api import service_api_ns
|
||||
from controllers.service_api.app.error import NotWorkflowAppError
|
||||
from controllers.service_api.schema import event_stream_response
|
||||
from controllers.service_api.wraps import FetchUserArg, WhereisUserArg, validate_app_token
|
||||
from core.app.apps.advanced_chat.app_generator import AdvancedChatAppGenerator
|
||||
from core.app.apps.base_app_generator import BaseAppGenerator
|
||||
@@ -44,6 +45,7 @@ register_response_schema_model(service_api_ns, EventStreamResponse)
|
||||
class WorkflowEventsApi(Resource):
|
||||
"""Service API for getting workflow execution events after resume."""
|
||||
|
||||
@event_stream_response(service_api_ns)
|
||||
@service_api_ns.doc("get_workflow_events")
|
||||
@service_api_ns.doc(description="Get workflow execution events stream after resume")
|
||||
@service_api_ns.doc(params={"task_id": "Workflow run ID"})
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
from typing import Any, Literal
|
||||
from typing import Any, Literal, override
|
||||
from uuid import UUID
|
||||
|
||||
from flask import request
|
||||
from pydantic import BaseModel, ConfigDict, Field, RootModel, field_validator, model_validator
|
||||
from pydantic import BaseModel, ConfigDict, Field, GetJsonSchemaHandler, RootModel, field_validator, model_validator
|
||||
from werkzeug.exceptions import Forbidden, NotFound
|
||||
|
||||
import services
|
||||
@@ -79,6 +79,13 @@ class DocumentStatusPayload(BaseModel):
|
||||
document_ids: list[str] = Field(default_factory=list, description="Document IDs to update")
|
||||
|
||||
|
||||
DOCUMENT_STATUS_ACTION_PARAM = {
|
||||
"description": "Action to perform: 'enable', 'disable', 'archive', or 'un_archive'",
|
||||
"enum": ["enable", "disable", "archive", "un_archive"],
|
||||
"type": "string",
|
||||
}
|
||||
|
||||
|
||||
class TagNamePayload(BaseModel):
|
||||
name: str = Field(..., min_length=1, max_length=50)
|
||||
|
||||
@@ -114,6 +121,45 @@ class TagUnbindingPayload(BaseModel):
|
||||
tag_id: str | None = None
|
||||
target_id: str
|
||||
|
||||
@classmethod
|
||||
@override
|
||||
def __get_pydantic_json_schema__(cls, _core_schema: object, _handler: GetJsonSchemaHandler) -> dict[str, object]:
|
||||
tag_id_property = {
|
||||
"description": "Legacy single tag ID accepted by the Service API.",
|
||||
"type": "string",
|
||||
}
|
||||
tag_ids_property = {
|
||||
"description": "Tag IDs to unbind. Use this for new integrations.",
|
||||
"items": {"type": "string"},
|
||||
"minItems": 1,
|
||||
"type": "array",
|
||||
}
|
||||
target_id_property = {"title": "Target Id", "type": "string"}
|
||||
return {
|
||||
"anyOf": [
|
||||
{
|
||||
"properties": {
|
||||
"tag_id": tag_id_property,
|
||||
"tag_ids": tag_ids_property,
|
||||
"target_id": target_id_property,
|
||||
},
|
||||
"required": ["tag_id", "target_id"],
|
||||
"type": "object",
|
||||
},
|
||||
{
|
||||
"properties": {
|
||||
"tag_id": {**tag_id_property, "nullable": True},
|
||||
"tag_ids": tag_ids_property,
|
||||
"target_id": target_id_property,
|
||||
},
|
||||
"required": ["tag_ids", "target_id"],
|
||||
"type": "object",
|
||||
},
|
||||
],
|
||||
"description": "Accepts either the legacy tag_id payload or the normalized tag_ids payload.",
|
||||
"title": cls.__name__,
|
||||
}
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def normalize_legacy_tag_id(cls, data: object) -> object:
|
||||
@@ -529,7 +575,7 @@ class DocumentStatusApi(DatasetApiResource):
|
||||
@service_api_ns.doc(
|
||||
params={
|
||||
"dataset_id": "Dataset ID",
|
||||
"action": "Action to perform: 'enable', 'disable', 'archive', or 'un_archive'",
|
||||
"action": DOCUMENT_STATUS_ACTION_PARAM,
|
||||
}
|
||||
)
|
||||
@service_api_ns.doc(
|
||||
|
||||
@@ -8,11 +8,12 @@ deprecated in generated API docs so clients migrate toward the canonical paths.
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from contextlib import ExitStack
|
||||
from typing import Any, Literal, Self
|
||||
from copy import deepcopy
|
||||
from typing import Any, Literal, Self, override
|
||||
from uuid import UUID
|
||||
|
||||
from flask import request, send_file
|
||||
from pydantic import BaseModel, Field, field_validator, model_validator
|
||||
from pydantic import BaseModel, Field, GetJsonSchemaHandler, field_validator, model_validator
|
||||
from sqlalchemy import desc, func, select
|
||||
from werkzeug.exceptions import Forbidden, NotFound
|
||||
|
||||
@@ -39,6 +40,7 @@ from controllers.service_api.dataset.error import (
|
||||
DocumentIndexingError,
|
||||
InvalidMetadataError,
|
||||
)
|
||||
from controllers.service_api.schema import binary_response
|
||||
from controllers.service_api.wraps import (
|
||||
DatasetApiResource,
|
||||
cloud_edition_billing_rate_limit_check,
|
||||
@@ -104,6 +106,36 @@ class DocumentTextUpdate(BaseModel):
|
||||
raise ValueError("Invalid doc_form.")
|
||||
return value
|
||||
|
||||
@classmethod
|
||||
@override
|
||||
def __get_pydantic_json_schema__(cls, core_schema: Any, handler: GetJsonSchemaHandler) -> dict[str, Any]:
|
||||
schema = handler.resolve_ref_schema(handler(core_schema))
|
||||
properties = schema.get("properties")
|
||||
if not isinstance(properties, dict):
|
||||
return schema
|
||||
|
||||
text_branch_properties = deepcopy(properties)
|
||||
text_branch_properties["text"] = _non_null_property_schema(properties.get("text"))
|
||||
text_branch_properties["name"] = _non_null_property_schema(properties.get("name"))
|
||||
|
||||
no_text_branch_properties = deepcopy(properties)
|
||||
no_text_branch_properties["text"] = {"type": "null"}
|
||||
|
||||
return {
|
||||
**schema,
|
||||
"anyOf": [
|
||||
{
|
||||
"properties": text_branch_properties,
|
||||
"required": ["name", "text"],
|
||||
"type": "object",
|
||||
},
|
||||
{
|
||||
"properties": no_text_branch_properties,
|
||||
"type": "object",
|
||||
},
|
||||
],
|
||||
}
|
||||
|
||||
@model_validator(mode="after")
|
||||
def check_text_and_name(self) -> Self:
|
||||
if self.text is not None and self.name is None:
|
||||
@@ -111,6 +143,24 @@ class DocumentTextUpdate(BaseModel):
|
||||
return self
|
||||
|
||||
|
||||
def _non_null_property_schema(property_schema: object) -> dict[str, Any]:
|
||||
if not isinstance(property_schema, dict):
|
||||
return {}
|
||||
|
||||
any_of = property_schema.get("anyOf")
|
||||
if isinstance(any_of, list):
|
||||
non_null_candidates = [
|
||||
candidate for candidate in any_of if isinstance(candidate, dict) and candidate.get("type") != "null"
|
||||
]
|
||||
if len(non_null_candidates) == 1:
|
||||
return {
|
||||
**{key: value for key, value in property_schema.items() if key != "anyOf"},
|
||||
**deepcopy(non_null_candidates[0]),
|
||||
}
|
||||
|
||||
return deepcopy(property_schema)
|
||||
|
||||
|
||||
class DocumentListQuery(BaseModel):
|
||||
page: int = Field(default=1, description="Page number")
|
||||
limit: int = Field(default=20, description="Number of items per page")
|
||||
@@ -463,8 +513,17 @@ class DeprecatedDocumentUpdateByTextApi(DatasetApiResource):
|
||||
|
||||
@service_api_ns.route(
|
||||
"/datasets/<uuid:dataset_id>/document/create_by_file",
|
||||
"/datasets/<uuid:dataset_id>/document/create-by-file",
|
||||
doc={
|
||||
"post": {
|
||||
"deprecated": True,
|
||||
"description": (
|
||||
"Deprecated legacy alias for creating a new document by uploading a file. "
|
||||
"Use /datasets/{dataset_id}/document/create-by-file instead."
|
||||
),
|
||||
}
|
||||
},
|
||||
)
|
||||
@service_api_ns.route("/datasets/<uuid:dataset_id>/document/create-by-file")
|
||||
class DocumentAddByFileApi(DatasetApiResource):
|
||||
"""Resource for documents."""
|
||||
|
||||
@@ -746,6 +805,7 @@ class DocumentListApi(DatasetApiResource):
|
||||
class DocumentBatchDownloadZipApi(DatasetApiResource):
|
||||
"""Download multiple uploaded-file documents as a single ZIP archive."""
|
||||
|
||||
@binary_response(service_api_ns, "application/zip")
|
||||
@service_api_ns.expect(service_api_ns.models[DocumentBatchDownloadZipPayload.__name__])
|
||||
@service_api_ns.doc("download_documents_as_zip")
|
||||
@service_api_ns.doc(description="Download selected uploaded documents as a single ZIP archive")
|
||||
@@ -758,11 +818,7 @@ class DocumentBatchDownloadZipApi(DatasetApiResource):
|
||||
404: "Document or dataset not found",
|
||||
}
|
||||
)
|
||||
@service_api_ns.response(
|
||||
200,
|
||||
"ZIP archive generated successfully",
|
||||
service_api_ns.models[BinaryFileResponse.__name__],
|
||||
)
|
||||
@service_api_ns.response(200, "ZIP archive generated successfully")
|
||||
@cloud_edition_billing_rate_limit_check("knowledge", "dataset")
|
||||
def post(self, tenant_id, dataset_id: UUID):
|
||||
payload = DocumentBatchDownloadZipPayload.model_validate(service_api_ns.payload or {})
|
||||
|
||||
@@ -24,6 +24,12 @@ from services.entities.knowledge_entities.knowledge_entities import (
|
||||
)
|
||||
from services.metadata_service import MetadataService
|
||||
|
||||
BUILT_IN_METADATA_ACTION_PARAM = {
|
||||
"description": "Action to perform: 'enable' or 'disable'",
|
||||
"enum": ["enable", "disable"],
|
||||
"type": "string",
|
||||
}
|
||||
|
||||
register_schema_model(service_api_ns, MetadataUpdatePayload)
|
||||
register_schema_models(
|
||||
service_api_ns,
|
||||
@@ -175,7 +181,7 @@ class DatasetMetadataBuiltInFieldServiceApi(DatasetApiResource):
|
||||
class DatasetMetadataBuiltInFieldActionServiceApi(DatasetApiResource):
|
||||
@service_api_ns.doc("toggle_built_in_field")
|
||||
@service_api_ns.doc(description="Enable or disable built-in metadata field")
|
||||
@service_api_ns.doc(params={"dataset_id": "Dataset ID", "action": "Action to perform: 'enable' or 'disable'"})
|
||||
@service_api_ns.doc(params={"dataset_id": "Dataset ID", "action": BUILT_IN_METADATA_ACTION_PARAM})
|
||||
@service_api_ns.doc(
|
||||
responses={
|
||||
200: "Action completed successfully",
|
||||
|
||||
@@ -19,6 +19,11 @@ from controllers.common.schema import (
|
||||
from controllers.service_api import service_api_ns
|
||||
from controllers.service_api.dataset.error import PipelineRunError
|
||||
from controllers.service_api.dataset.rag_pipeline.serializers import serialize_upload_file
|
||||
from controllers.service_api.schema import (
|
||||
event_stream_response,
|
||||
json_or_event_stream_response,
|
||||
multipart_file_params,
|
||||
)
|
||||
from controllers.service_api.wraps import DatasetApiResource
|
||||
from core.app.apps.pipeline.pipeline_generator import PipelineGenerator
|
||||
from core.app.entities.app_invoke_entities import InvokeFrom
|
||||
@@ -137,6 +142,7 @@ class DatasourcePluginsApi(DatasetApiResource):
|
||||
class DatasourceNodeRunApi(DatasetApiResource):
|
||||
"""Resource for datasource node run."""
|
||||
|
||||
@event_stream_response(service_api_ns)
|
||||
@service_api_ns.doc(shortcut="pipeline_datasource_node_run")
|
||||
@service_api_ns.doc(description="Run a datasource node for a rag pipeline")
|
||||
@service_api_ns.doc(
|
||||
@@ -195,6 +201,7 @@ class DatasourceNodeRunApi(DatasetApiResource):
|
||||
class PipelineRunApi(DatasetApiResource):
|
||||
"""Resource for datasource node run."""
|
||||
|
||||
@json_or_event_stream_response(service_api_ns)
|
||||
@service_api_ns.doc(shortcut="pipeline_datasource_node_run")
|
||||
@service_api_ns.doc(description="Run a datasource node for a rag pipeline")
|
||||
@service_api_ns.doc(
|
||||
@@ -250,6 +257,7 @@ class KnowledgebasePipelineFileUploadApi(DatasetApiResource):
|
||||
|
||||
@service_api_ns.doc(shortcut="knowledgebase_pipeline_file_upload")
|
||||
@service_api_ns.doc(description="Upload a file to a knowledgebase pipeline")
|
||||
@service_api_ns.doc(consumes=["multipart/form-data"], params=multipart_file_params(include_user=False))
|
||||
@service_api_ns.doc(
|
||||
responses={
|
||||
201: "File uploaded successfully",
|
||||
|
||||
@@ -0,0 +1,113 @@
|
||||
"""Service API OpenAPI documentation helpers.
|
||||
|
||||
These helpers keep documentation-only request shapes next to controller
|
||||
definitions without changing the Pydantic models used for runtime validation.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
from copy import deepcopy
|
||||
from typing import cast
|
||||
|
||||
from flask_restx import Namespace
|
||||
from pydantic import BaseModel
|
||||
|
||||
USER_PROPERTY_SCHEMA: dict[str, object] = {"description": "End user identifier", "type": "string"}
|
||||
USER_QUERY_PARAM: dict[str, object] = {"description": "End user identifier", "in": "query", "type": "string"}
|
||||
USER_FORM_PARAM: dict[str, object] = {"description": "End user identifier", "in": "formData", "type": "string"}
|
||||
FILE_FORM_PARAM: dict[str, object] = {"in": "formData", "required": True, "type": "file"}
|
||||
USER_FETCH_FROM_ATTR = "_dify_service_api_user_fetch_from"
|
||||
USER_REQUIRED_ATTR = "_dify_service_api_user_required"
|
||||
JSON_USER_FETCH_FROM = "JSON"
|
||||
|
||||
|
||||
def expect_with_user(namespace: Namespace, model: type[BaseModel]):
|
||||
"""Document a JSON request body as ``model`` plus Service API ``user``."""
|
||||
|
||||
source_model = namespace.models[model.__name__]
|
||||
model_name = f"{model.__name__}WithUser"
|
||||
|
||||
def decorator(view_func):
|
||||
required = _json_user_required(view_func)
|
||||
schema = cast(dict[str, object], deepcopy(source_model.__schema__))
|
||||
_add_user_property(schema, required=required)
|
||||
if model_name not in namespace.models:
|
||||
namespace.schema_model(model_name, schema)
|
||||
return namespace.expect(namespace.models[model_name], validate=False)(view_func)
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
def expect_user_json(namespace: Namespace):
|
||||
"""Document a JSON request body that only carries the Service API ``user``."""
|
||||
|
||||
def decorator(view_func):
|
||||
required = _json_user_required(view_func)
|
||||
schema: dict[str, object] = {"properties": {}, "title": "ServiceApiUserPayload", "type": "object"}
|
||||
_add_user_property(schema, required=required)
|
||||
model_name = "RequiredServiceApiUserPayload" if required else "OptionalServiceApiUserPayload"
|
||||
if model_name not in namespace.models:
|
||||
namespace.schema_model(model_name, schema)
|
||||
return namespace.expect(namespace.models[model_name], validate=False)(view_func)
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
def multipart_file_params(*, include_user: bool) -> dict[str, dict[str, object]]:
|
||||
params: dict[str, dict[str, object]] = {"file": FILE_FORM_PARAM}
|
||||
if include_user:
|
||||
params["user"] = USER_FORM_PARAM
|
||||
return deepcopy(params)
|
||||
|
||||
|
||||
def json_or_event_stream_response(namespace: Namespace):
|
||||
return namespace.doc(produces=["application/json", "text/event-stream"])
|
||||
|
||||
|
||||
def event_stream_response(namespace: Namespace):
|
||||
return namespace.doc(produces=["text/event-stream"])
|
||||
|
||||
|
||||
def binary_response(namespace: Namespace, media_type: str | Sequence[str]):
|
||||
media_types = [media_type] if isinstance(media_type, str) else list(media_type)
|
||||
return namespace.doc(produces=media_types)
|
||||
|
||||
|
||||
def _json_user_required(view_func) -> bool:
|
||||
fetch_from = getattr(view_func, USER_FETCH_FROM_ATTR, None)
|
||||
if fetch_from != JSON_USER_FETCH_FROM:
|
||||
raise ValueError("JSON user documentation must match validate_app_token(fetch_user_arg=WhereisUserArg.JSON)")
|
||||
|
||||
return bool(getattr(view_func, USER_REQUIRED_ATTR, False))
|
||||
|
||||
|
||||
def _add_user_property(schema: dict[str, object], *, required: bool) -> None:
|
||||
variants: list[dict[str, object]] = []
|
||||
for keyword in ("anyOf", "oneOf"):
|
||||
candidates = schema.get(keyword)
|
||||
if isinstance(candidates, list):
|
||||
variants.extend(candidate for candidate in candidates if isinstance(candidate, dict))
|
||||
|
||||
if variants:
|
||||
for variant in variants:
|
||||
_add_user_property_to_object_schema(variant, required=required)
|
||||
|
||||
_add_user_property_to_object_schema(schema, required=required)
|
||||
|
||||
|
||||
def _add_user_property_to_object_schema(schema: dict[str, object], *, required: bool) -> None:
|
||||
properties = schema.setdefault("properties", {})
|
||||
if isinstance(properties, dict):
|
||||
cast(dict[str, object], properties)["user"] = USER_PROPERTY_SCHEMA
|
||||
|
||||
if required:
|
||||
required_fields = schema.setdefault("required", [])
|
||||
if isinstance(required_fields, list) and "user" not in required_fields:
|
||||
required_fields.append("user")
|
||||
else:
|
||||
required_fields = schema.get("required")
|
||||
if isinstance(required_fields, list) and "user" in required_fields:
|
||||
required_fields.remove("user")
|
||||
if required_fields == []:
|
||||
schema.pop("required", None)
|
||||
@@ -4,16 +4,23 @@ import time
|
||||
from collections.abc import Callable
|
||||
from enum import StrEnum, auto
|
||||
from functools import wraps
|
||||
from typing import cast, overload
|
||||
from typing import Protocol, cast, overload
|
||||
|
||||
from flask import current_app, request
|
||||
from flask_login import user_logged_in
|
||||
from flask_restx import Resource
|
||||
from flask_restx.utils import merge
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy import select
|
||||
from werkzeug.exceptions import Forbidden, NotFound, Unauthorized
|
||||
|
||||
from configs import dify_config
|
||||
from controllers.service_api.schema import (
|
||||
USER_FETCH_FROM_ATTR,
|
||||
USER_FORM_PARAM,
|
||||
USER_QUERY_PARAM,
|
||||
USER_REQUIRED_ATTR,
|
||||
)
|
||||
from enums.cloud_plan import CloudPlan
|
||||
from extensions.ext_database import db
|
||||
from extensions.ext_redis import redis_client
|
||||
@@ -28,6 +35,12 @@ from services.feature_service import FeatureService
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class _RestxDocumentedView(Protocol):
|
||||
"""Callable view object carrying Flask-RESTX documentation metadata."""
|
||||
|
||||
__apidoc__: dict[str, object]
|
||||
|
||||
|
||||
class WhereisUserArg(StrEnum):
|
||||
"""
|
||||
Enum for whereis_user_arg.
|
||||
@@ -43,6 +56,35 @@ class FetchUserArg(BaseModel):
|
||||
required: bool = False
|
||||
|
||||
|
||||
APP_TOKEN_FORBIDDEN_RESPONSE = {
|
||||
403: "Forbidden - token scope, app, dataset, or workspace access denied",
|
||||
}
|
||||
|
||||
DATASET_TOKEN_AUTH_RESPONSES = {
|
||||
401: "Unauthorized - invalid API token",
|
||||
403: "Forbidden - dataset API access or workspace access denied",
|
||||
}
|
||||
|
||||
|
||||
def _document_app_token_contract(view_func: Callable[..., object], fetch_user_arg: FetchUserArg | None) -> None:
|
||||
doc: dict[str, object] = {"responses": APP_TOKEN_FORBIDDEN_RESPONSE}
|
||||
if fetch_user_arg is not None:
|
||||
setattr(view_func, USER_FETCH_FROM_ATTR, fetch_user_arg.fetch_from.name)
|
||||
setattr(view_func, USER_REQUIRED_ATTR, fetch_user_arg.required)
|
||||
match fetch_user_arg.fetch_from:
|
||||
case WhereisUserArg.QUERY:
|
||||
doc["params"] = {"user": {**USER_QUERY_PARAM, "required": fetch_user_arg.required}}
|
||||
case WhereisUserArg.FORM:
|
||||
doc["params"] = {"user": {**USER_FORM_PARAM, "required": fetch_user_arg.required}}
|
||||
case WhereisUserArg.JSON:
|
||||
pass
|
||||
|
||||
cast(_RestxDocumentedView, view_func).__apidoc__ = cast(
|
||||
dict[str, object],
|
||||
merge(getattr(view_func, "__apidoc__", {}), doc),
|
||||
)
|
||||
|
||||
|
||||
@overload
|
||||
def validate_app_token[**P, R](view: Callable[P, R]) -> Callable[P, R]: ...
|
||||
|
||||
@@ -126,6 +168,7 @@ def validate_app_token[**P, R](
|
||||
|
||||
return view_func(*args, **kwargs)
|
||||
|
||||
_document_app_token_contract(decorated_view, fetch_user_arg)
|
||||
return decorated_view
|
||||
|
||||
if view is None:
|
||||
@@ -343,6 +386,8 @@ def validate_and_get_api_token(scope: str | None = None):
|
||||
|
||||
|
||||
class DatasetApiResource(Resource):
|
||||
__apidoc__ = {"responses": DATASET_TOKEN_AUTH_RESPONSES}
|
||||
|
||||
method_decorators = [validate_dataset_token]
|
||||
|
||||
def get_dataset(self, dataset_id: str, tenant_id: str) -> Dataset:
|
||||
|
||||
Reference in New Issue
Block a user