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:
Stephen Zhou
2026-06-17 14:25:30 +00:00
committed by GitHub
co-authored by autofix-ci[bot]
parent 9021b3f5be
commit baf775134e
51 changed files with 2778 additions and 1369 deletions
+42 -2
View File
@@ -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",
+5 -6
View File
@@ -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"})
+2
View File
@@ -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"})
+2 -1
View File
@@ -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")
+10 -2
View File
@@ -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"})
+49 -3
View File
@@ -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",
+113
View File
@@ -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)
+46 -1
View File
@@ -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: