mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
refactor(api): migrate web chat endpoints to BaseModel (#37962)
This commit is contained in:
@@ -1,13 +1,11 @@
|
||||
import logging
|
||||
|
||||
from flask import request
|
||||
from flask_restx import fields, marshal_with
|
||||
from pydantic import field_validator
|
||||
from werkzeug.exceptions import InternalServerError
|
||||
|
||||
import services
|
||||
from controllers.common.controller_schemas import TextToAudioPayload as TextToAudioPayloadBase
|
||||
from controllers.common.fields import AudioBinaryResponse, AudioTranscriptResponse
|
||||
from controllers.web import web_ns
|
||||
from controllers.web.error import (
|
||||
AppUnavailableError,
|
||||
@@ -23,8 +21,9 @@ from controllers.web.error import (
|
||||
from controllers.web.wraps import WebApiResource
|
||||
from core.errors.error import ModelCurrentlyNotSupportError, ProviderTokenNotInitError, QuotaExceededError
|
||||
from extensions.ext_database import db
|
||||
from fields.base import ResponseModel
|
||||
from graphon.model_runtime.errors.invoke import InvokeError
|
||||
from libs.helper import uuid_value
|
||||
from libs.helper import dump_response, uuid_value
|
||||
from models.model import App, EndUser
|
||||
from services.app_ref_service import AppRefService
|
||||
from services.audio_service import AudioService
|
||||
@@ -38,6 +37,10 @@ from services.errors.audio import (
|
||||
from ..common.schema import register_response_schema_models, register_schema_models
|
||||
|
||||
|
||||
class AudioToTextResponse(ResponseModel):
|
||||
text: str
|
||||
|
||||
|
||||
class TextToAudioPayload(TextToAudioPayloadBase):
|
||||
@field_validator("message_id")
|
||||
@classmethod
|
||||
@@ -48,18 +51,13 @@ class TextToAudioPayload(TextToAudioPayloadBase):
|
||||
|
||||
|
||||
register_schema_models(web_ns, TextToAudioPayload)
|
||||
register_response_schema_models(web_ns, AudioBinaryResponse, AudioTranscriptResponse)
|
||||
register_response_schema_models(web_ns, AudioToTextResponse)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@web_ns.route("/audio-to-text")
|
||||
class AudioApi(WebApiResource):
|
||||
audio_to_text_response_fields = {
|
||||
"text": fields.String,
|
||||
}
|
||||
|
||||
@marshal_with(audio_to_text_response_fields)
|
||||
@web_ns.doc("Audio to Text")
|
||||
@web_ns.doc(description="Convert audio file to text using speech-to-text service.")
|
||||
@web_ns.doc(
|
||||
@@ -73,7 +71,7 @@ class AudioApi(WebApiResource):
|
||||
500: "Internal Server Error",
|
||||
}
|
||||
)
|
||||
@web_ns.response(200, "Success", web_ns.models[AudioTranscriptResponse.__name__])
|
||||
@web_ns.response(200, "Success", web_ns.models[AudioToTextResponse.__name__])
|
||||
def post(self, app_model: App, end_user: EndUser):
|
||||
"""Convert audio to text"""
|
||||
file = request.files["file"]
|
||||
@@ -81,7 +79,7 @@ class AudioApi(WebApiResource):
|
||||
try:
|
||||
response = AudioService.transcript_asr(app_model=app_model, file=file, end_user=end_user.external_user_id)
|
||||
|
||||
return response
|
||||
return dump_response(AudioToTextResponse, response)
|
||||
except services.errors.app_model_config.AppModelConfigBrokenError:
|
||||
logger.exception("App model config broken.")
|
||||
raise AppUnavailableError()
|
||||
@@ -122,7 +120,8 @@ class TextApi(WebApiResource):
|
||||
500: "Internal Server Error",
|
||||
}
|
||||
)
|
||||
@web_ns.response(200, "Success", web_ns.models[AudioBinaryResponse.__name__])
|
||||
# response-contract:ignore provider audio bytes; TODO: model binary audio response if shape is standardized.
|
||||
@web_ns.response(200, "Success")
|
||||
def post(self, app_model: App, end_user: EndUser):
|
||||
"""Convert text to audio"""
|
||||
try:
|
||||
@@ -139,7 +138,7 @@ class TextApi(WebApiResource):
|
||||
message_id,
|
||||
end_user_id=end_user.id,
|
||||
)
|
||||
response = AudioService.transcript_tts(
|
||||
return AudioService.transcript_tts(
|
||||
app_model=app_model,
|
||||
session=db.session(),
|
||||
text=text,
|
||||
@@ -147,8 +146,6 @@ class TextApi(WebApiResource):
|
||||
end_user=end_user.external_user_id,
|
||||
message_ref=message_ref,
|
||||
)
|
||||
|
||||
return response
|
||||
except services.errors.app_model_config.AppModelConfigBrokenError:
|
||||
logger.exception("App model config broken.")
|
||||
raise AppUnavailableError()
|
||||
|
||||
@@ -133,6 +133,7 @@ class CompletionApi(WebApiResource):
|
||||
streaming=streaming,
|
||||
)
|
||||
|
||||
# response-contract:ignore compact_generate_response
|
||||
return helper.compact_generate_response(response)
|
||||
except services.errors.conversation.ConversationNotExistsError:
|
||||
raise NotFound("Conversation Not Exists.")
|
||||
@@ -185,7 +186,7 @@ class CompletionStopApi(WebApiResource):
|
||||
app_mode=AppMode.value_of(app_model.mode),
|
||||
)
|
||||
|
||||
return {"result": "success"}, 200
|
||||
return SimpleResultResponse(result="success").model_dump(mode="json"), 200
|
||||
|
||||
|
||||
@web_ns.route("/chat-messages")
|
||||
@@ -235,6 +236,7 @@ class ChatApi(WebApiResource):
|
||||
streaming=streaming,
|
||||
)
|
||||
|
||||
# response-contract:ignore compact_generate_response
|
||||
return helper.compact_generate_response(response)
|
||||
except services.errors.conversation.ConversationNotExistsError:
|
||||
raise NotFound("Conversation Not Exists.")
|
||||
@@ -290,4 +292,4 @@ class ChatStopApi(WebApiResource):
|
||||
app_mode=app_mode,
|
||||
)
|
||||
|
||||
return {"result": "success"}, 200
|
||||
return SimpleResultResponse(result="success").model_dump(mode="json"), 200
|
||||
|
||||
@@ -13,6 +13,7 @@ from controllers.web import web_ns
|
||||
from controllers.web.wraps import WebApiResource
|
||||
from extensions.ext_database import db
|
||||
from fields.file_fields import FileResponse
|
||||
from libs.helper import dump_response
|
||||
from models.model import App, EndUser
|
||||
from services.file_service import FileService
|
||||
|
||||
@@ -84,5 +85,4 @@ class FileApi(WebApiResource):
|
||||
except services.errors.file.UnsupportedFileTypeError:
|
||||
raise UnsupportedFileTypeError()
|
||||
|
||||
response = FileResponse.model_validate(upload_file, from_attributes=True)
|
||||
return response.model_dump(mode="json"), 201
|
||||
return dump_response(FileResponse, upload_file), 201
|
||||
|
||||
@@ -29,6 +29,7 @@ from extensions.ext_database import db
|
||||
from fields.file_fields import FileResponse, FileWithSignedUrl
|
||||
from graphon.file import helpers as file_helpers
|
||||
from libs.exception import BaseHTTPException
|
||||
from libs.helper import dump_response
|
||||
from repositories.factory import DifyAPIRepositoryFactory
|
||||
from services.file_service import FileService
|
||||
from services.human_input_file_upload_service import (
|
||||
@@ -141,8 +142,7 @@ def _upload_local_file(context):
|
||||
except services.errors.file.BlockedFileExtensionError as exc:
|
||||
raise BlockedFileExtensionError() from exc
|
||||
|
||||
response = FileResponse.model_validate(upload_file, from_attributes=True)
|
||||
return upload_file.id, response
|
||||
return upload_file.id, dump_response(FileResponse, upload_file)
|
||||
|
||||
|
||||
def _upload_remote_file(context, url: str):
|
||||
@@ -186,7 +186,7 @@ def _upload_remote_file(context, url: str):
|
||||
created_by=upload_file.created_by,
|
||||
created_at=int(upload_file.created_at.timestamp()),
|
||||
)
|
||||
return upload_file.id, response
|
||||
return upload_file.id, response.model_dump(mode="json")
|
||||
|
||||
|
||||
@web_ns.route("/human-input-forms/files")
|
||||
@@ -209,4 +209,5 @@ class HumanInputFileUploadApi(Resource):
|
||||
file_id, response = _upload_local_file(context=context)
|
||||
|
||||
upload_service.record_upload_file(context=context, file_id=file_id)
|
||||
return response.model_dump(mode="json"), 201
|
||||
# response-contract:ignore pre-dumped response. See above
|
||||
return response, 201
|
||||
|
||||
@@ -2,14 +2,12 @@
|
||||
Web App Human Input Form APIs.
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
from collections.abc import Sequence
|
||||
from typing import Any, NotRequired, TypedDict
|
||||
from typing import Self
|
||||
|
||||
from flask import Response, request
|
||||
from flask import request
|
||||
from flask_restx import Resource
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
from werkzeug.exceptions import Forbidden
|
||||
@@ -20,35 +18,58 @@ from controllers.common.human_input import HumanInputFormSubmitPayload, stringif
|
||||
from controllers.common.schema import register_response_schema_models, register_schema_models
|
||||
from controllers.web import web_ns
|
||||
from controllers.web.error import WebFormRateLimitExceededError
|
||||
from controllers.web.site import serialize_app_site_payload
|
||||
from core.workflow.nodes.human_input.entities import FormInputConfig
|
||||
from controllers.web.site import WebAppSiteResponse
|
||||
from core.workflow.nodes.human_input.entities import FormInputConfig, UserActionConfig
|
||||
from extensions.ext_database import db
|
||||
from libs.helper import RateLimiter, extract_remote_ip, to_timestamp
|
||||
from fields.base import ResponseModel
|
||||
from libs.helper import RateLimiter, dump_response, extract_remote_ip, to_timestamp
|
||||
from models.account import TenantStatus
|
||||
from models.model import App, Site
|
||||
from repositories.factory import DifyAPIRepositoryFactory
|
||||
from services.feature_service import FeatureService
|
||||
from services.human_input_file_upload_service import HumanInputFileUploadService
|
||||
from services.human_input_service import Form, FormNotFoundError, HumanInputService
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class HumanInputUploadTokenResponse(BaseModel):
|
||||
class HumanInputUploadTokenResponse(ResponseModel):
|
||||
upload_token: str
|
||||
expires_at: int
|
||||
|
||||
|
||||
class HumanInputFormDefinitionResponse(BaseModel):
|
||||
form_content: Any
|
||||
inputs: Any
|
||||
class HumanInputFormDefinitionResponse(ResponseModel):
|
||||
form_content: str
|
||||
inputs: list[FormInputConfig]
|
||||
resolved_default_values: dict[str, str]
|
||||
user_actions: Any
|
||||
user_actions: list[UserActionConfig]
|
||||
expiration_time: int
|
||||
site: dict[str, Any] | None = Field(default=None)
|
||||
site: WebAppSiteResponse | None = None
|
||||
|
||||
@classmethod
|
||||
def from_form(
|
||||
cls,
|
||||
form: Form,
|
||||
*,
|
||||
inputs: Sequence[FormInputConfig] = (),
|
||||
site: WebAppSiteResponse | None = None,
|
||||
) -> Self:
|
||||
definition_payload = form.get_definition().model_dump(mode="json")
|
||||
expiration_time = to_timestamp(form.expiration_time)
|
||||
if expiration_time is None:
|
||||
raise ValueError("Human input form expiration_time is required")
|
||||
return cls(
|
||||
form_content=definition_payload["rendered_content"],
|
||||
inputs=list(inputs),
|
||||
resolved_default_values=stringify_form_default_values(definition_payload["default_values"]),
|
||||
user_actions=definition_payload["user_actions"],
|
||||
expiration_time=expiration_time,
|
||||
site=site,
|
||||
)
|
||||
|
||||
|
||||
class HumanInputFormSubmitResponse(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
class HumanInputFormSubmitResponse(ResponseModel):
|
||||
pass
|
||||
|
||||
|
||||
register_schema_models(web_ns, HumanInputFormSubmitPayload)
|
||||
@@ -86,40 +107,26 @@ def _create_upload_service() -> HumanInputFileUploadService:
|
||||
)
|
||||
|
||||
|
||||
class FormDefinitionPayload(TypedDict):
|
||||
form_content: Any
|
||||
inputs: Any
|
||||
resolved_default_values: dict[str, str]
|
||||
user_actions: Any
|
||||
expiration_time: int
|
||||
site: NotRequired[dict]
|
||||
|
||||
|
||||
def _jsonify_form_definition(
|
||||
form: Form,
|
||||
*,
|
||||
inputs: Sequence[FormInputConfig] = (),
|
||||
site_payload: dict | None = None,
|
||||
) -> Response:
|
||||
"""Return the form payload (optionally with site) as a JSON response."""
|
||||
definition_payload = form.get_definition().model_dump(mode="json")
|
||||
payload: FormDefinitionPayload = {
|
||||
"form_content": definition_payload["rendered_content"],
|
||||
"inputs": [i.model_dump(mode="json") for i in inputs],
|
||||
"resolved_default_values": stringify_form_default_values(definition_payload["default_values"]),
|
||||
"user_actions": definition_payload["user_actions"],
|
||||
"expiration_time": to_timestamp(form.expiration_time),
|
||||
}
|
||||
if site_payload is not None:
|
||||
payload["site"] = site_payload
|
||||
return Response(json.dumps(payload, ensure_ascii=False), mimetype="application/json")
|
||||
|
||||
|
||||
@web_ns.route("/form/human_input/<string:form_token>/upload-token")
|
||||
class HumanInputFormUploadTokenApi(Resource):
|
||||
"""API for issuing HITL upload tokens for active human input forms."""
|
||||
|
||||
@web_ns.response(200, "Success", web_ns.models[HumanInputUploadTokenResponse.__name__])
|
||||
@web_ns.doc("create_human_input_form_upload_token")
|
||||
@web_ns.doc(description="Issue an upload token for an active human input form")
|
||||
@web_ns.doc(params={"form_token": "Human input form token"})
|
||||
@web_ns.doc(
|
||||
responses={
|
||||
200: "Upload token issued successfully",
|
||||
404: "Form not found",
|
||||
412: "Form already submitted or expired",
|
||||
429: "Too many requests",
|
||||
}
|
||||
)
|
||||
@web_ns.response(
|
||||
200,
|
||||
"Upload token issued successfully",
|
||||
web_ns.models[HumanInputUploadTokenResponse.__name__],
|
||||
)
|
||||
def post(self, form_token: str):
|
||||
"""
|
||||
Issue an upload token for a human input form.
|
||||
@@ -136,11 +143,9 @@ class HumanInputFormUploadTokenApi(Resource):
|
||||
except FormNotFoundError:
|
||||
raise NotFoundError("Form not found")
|
||||
|
||||
response = HumanInputUploadTokenResponse(
|
||||
upload_token=token.upload_token,
|
||||
expires_at=to_timestamp(token.expires_at),
|
||||
)
|
||||
return response.model_dump(mode="json"), 200
|
||||
return HumanInputUploadTokenResponse(
|
||||
upload_token=token.upload_token, expires_at=to_timestamp(token.expires_at)
|
||||
).model_dump(mode="json"), 200
|
||||
|
||||
|
||||
@web_ns.route("/form/human_input/<string:form_token>")
|
||||
@@ -150,7 +155,23 @@ class HumanInputFormApi(Resource):
|
||||
# NOTE(QuantumGhost): this endpoint is unauthenticated on purpose for now.
|
||||
|
||||
# def get(self, _app_model: App, _end_user: EndUser, form_token: str):
|
||||
@web_ns.response(200, "Success", web_ns.models[HumanInputFormDefinitionResponse.__name__])
|
||||
@web_ns.doc("get_human_input_form")
|
||||
@web_ns.doc(description="Get a human input form definition by token")
|
||||
@web_ns.doc(params={"form_token": "Human input form token"})
|
||||
@web_ns.doc(
|
||||
responses={
|
||||
200: "Form retrieved successfully",
|
||||
403: "Forbidden",
|
||||
404: "Form not found",
|
||||
412: "Form already submitted or expired",
|
||||
429: "Too many requests",
|
||||
}
|
||||
)
|
||||
@web_ns.response(
|
||||
200,
|
||||
"Form retrieved successfully",
|
||||
web_ns.models[HumanInputFormDefinitionResponse.__name__],
|
||||
)
|
||||
def get(self, form_token: str):
|
||||
"""
|
||||
Get human input form definition by token.
|
||||
@@ -172,17 +193,47 @@ class HumanInputFormApi(Resource):
|
||||
|
||||
service.ensure_form_active(form)
|
||||
app_model, site = _get_app_site_from_form(form)
|
||||
tenant = app_model.tenant
|
||||
if tenant is None:
|
||||
raise Forbidden()
|
||||
inputs = service.resolve_form_inputs(form)
|
||||
features = FeatureService.get_features(app_model.tenant_id, exclude_vector_space=True)
|
||||
|
||||
return _jsonify_form_definition(
|
||||
form,
|
||||
inputs=inputs,
|
||||
site_payload=serialize_app_site_payload(app_model, site, None),
|
||||
return dump_response(
|
||||
HumanInputFormDefinitionResponse,
|
||||
HumanInputFormDefinitionResponse.from_form(
|
||||
form,
|
||||
inputs=inputs,
|
||||
site=WebAppSiteResponse.from_app_site(
|
||||
tenant=tenant,
|
||||
app_model=app_model,
|
||||
site=site,
|
||||
end_user_id=None,
|
||||
features=features,
|
||||
can_replace_logo=features.can_replace_logo,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
# def post(self, _app_model: App, _end_user: EndUser, form_token: str):
|
||||
@web_ns.response(200, "Success", web_ns.models[HumanInputFormSubmitResponse.__name__])
|
||||
@web_ns.expect(web_ns.models[HumanInputFormSubmitPayload.__name__])
|
||||
@web_ns.doc("submit_human_input_form")
|
||||
@web_ns.doc(description="Submit a human input form by token")
|
||||
@web_ns.doc(params={"form_token": "Human input form token"})
|
||||
@web_ns.doc(
|
||||
responses={
|
||||
200: "Form submitted successfully",
|
||||
400: "Bad request - invalid submission data",
|
||||
404: "Form not found",
|
||||
412: "Form already submitted or expired",
|
||||
429: "Too many requests",
|
||||
}
|
||||
)
|
||||
@web_ns.response(
|
||||
200,
|
||||
"Form submitted successfully",
|
||||
web_ns.models[HumanInputFormSubmitResponse.__name__],
|
||||
)
|
||||
def post(self, form_token: str):
|
||||
"""
|
||||
Submit human input form by token.
|
||||
@@ -225,7 +276,7 @@ class HumanInputFormApi(Resource):
|
||||
except FormNotFoundError:
|
||||
raise NotFoundError("Form not found")
|
||||
|
||||
return {}, 200
|
||||
return HumanInputFormSubmitResponse().model_dump(mode="json"), 200
|
||||
|
||||
|
||||
def _get_app_site_from_form(form: Form) -> tuple[App, Site]:
|
||||
@@ -238,7 +289,7 @@ def _get_app_site_from_form(form: Form) -> tuple[App, Site]:
|
||||
if site is None:
|
||||
raise Forbidden()
|
||||
|
||||
if app_model.tenant and app_model.tenant.status == TenantStatus.ARCHIVE:
|
||||
if app_model.tenant is None or app_model.tenant.status == TenantStatus.ARCHIVE:
|
||||
raise Forbidden()
|
||||
|
||||
return app_model, site
|
||||
|
||||
@@ -188,6 +188,7 @@ class MessageMoreLikeThisApi(WebApiResource):
|
||||
streaming=streaming,
|
||||
)
|
||||
|
||||
# response-contract:ignore compact_generate_response
|
||||
return helper.compact_generate_response(response)
|
||||
except MessageNotExistsError:
|
||||
raise NotFound("Message Not Exists.")
|
||||
|
||||
@@ -65,11 +65,10 @@ class RemoteFileInfoApi(WebApiResource):
|
||||
# failed back to get method
|
||||
resp = remote_fetcher.make_request("GET", decoded_url, timeout=3)
|
||||
resp.raise_for_status()
|
||||
info = RemoteFileInfo(
|
||||
return RemoteFileInfo(
|
||||
file_type=resp.headers.get("Content-Type", "application/octet-stream"),
|
||||
file_length=int(resp.headers.get("Content-Length", -1)),
|
||||
)
|
||||
return info.model_dump(mode="json")
|
||||
).model_dump(mode="json")
|
||||
|
||||
|
||||
@web_ns.route("/remote-files/upload")
|
||||
@@ -141,7 +140,7 @@ class RemoteFileUploadApi(WebApiResource):
|
||||
except services.errors.file.UnsupportedFileTypeError:
|
||||
raise UnsupportedFileTypeError
|
||||
|
||||
payload1 = FileWithSignedUrl(
|
||||
return FileWithSignedUrl(
|
||||
id=upload_file.id,
|
||||
name=upload_file.name,
|
||||
size=upload_file.size,
|
||||
@@ -150,5 +149,4 @@ class RemoteFileUploadApi(WebApiResource):
|
||||
mime_type=upload_file.mime_type,
|
||||
created_by=upload_file.created_by,
|
||||
created_at=int(upload_file.created_at.timestamp()),
|
||||
)
|
||||
return payload1.model_dump(mode="json"), 201
|
||||
).model_dump(mode="json"), 201
|
||||
|
||||
@@ -49,9 +49,7 @@ class SavedMessageListApi(WebApiResource):
|
||||
adapter = TypeAdapter(SavedMessageItem)
|
||||
items = [adapter.validate_python(message, from_attributes=True) for message in pagination.data]
|
||||
return SavedMessageInfiniteScrollPagination(
|
||||
limit=pagination.limit,
|
||||
has_more=pagination.has_more,
|
||||
data=items,
|
||||
limit=pagination.limit, has_more=pagination.has_more, data=items
|
||||
).model_dump(mode="json")
|
||||
|
||||
@web_ns.doc("Save Message")
|
||||
@@ -102,6 +100,7 @@ class SavedMessageApi(WebApiResource):
|
||||
500: "Internal Server Error",
|
||||
}
|
||||
)
|
||||
@web_ns.response(204, "Message removed successfully")
|
||||
def delete(self, app_model: App, end_user: EndUser, message_id: UUID):
|
||||
message_id_str = str(message_id)
|
||||
|
||||
|
||||
+100
-132
@@ -1,7 +1,6 @@
|
||||
from typing import Any, cast
|
||||
from typing import Any, Self
|
||||
|
||||
from flask_restx import fields, marshal, marshal_with
|
||||
from pydantic import Field
|
||||
from pydantic import AliasChoices, Field, computed_field
|
||||
from sqlalchemy import select
|
||||
from werkzeug.exceptions import Forbidden
|
||||
|
||||
@@ -11,30 +10,19 @@ from controllers.web import web_ns
|
||||
from controllers.web.wraps import WebApiResource
|
||||
from extensions.ext_database import db
|
||||
from fields.base import ResponseModel
|
||||
from libs.helper import AppIconUrlField
|
||||
from models.account import TenantStatus
|
||||
from libs.helper import build_icon_url
|
||||
from models.account import Tenant, TenantStatus
|
||||
from models.model import App, EndUser, Site
|
||||
from services.feature_service import FeatureModel, FeatureService
|
||||
|
||||
|
||||
class AppSiteModelConfigResponse(ResponseModel):
|
||||
opening_statement: str | None = None
|
||||
suggested_questions: Any
|
||||
suggested_questions_after_answer: Any
|
||||
more_like_this: Any
|
||||
model: Any
|
||||
user_input_form: Any
|
||||
pre_prompt: str | None = None
|
||||
|
||||
|
||||
class AppSiteResponse(ResponseModel):
|
||||
title: str | None = None
|
||||
class WebSiteResponse(ResponseModel):
|
||||
title: str
|
||||
chat_color_theme: str | None = None
|
||||
chat_color_theme_inverted: bool | None = None
|
||||
chat_color_theme_inverted: bool
|
||||
icon_type: str | None = None
|
||||
icon: str | None = None
|
||||
icon_background: str | None = None
|
||||
icon_url: str | None = None
|
||||
description: str | None = None
|
||||
copyright: str | None = None
|
||||
privacy_policy: str | None = None
|
||||
@@ -45,65 +33,98 @@ class AppSiteResponse(ResponseModel):
|
||||
show_workflow_steps: bool | None = None
|
||||
use_icon_as_answer_icon: bool | None = None
|
||||
|
||||
@computed_field(return_type=str | None) # type: ignore[prop-decorator]
|
||||
@property
|
||||
def icon_url(self) -> str | None:
|
||||
return build_icon_url(self.icon_type, self.icon)
|
||||
|
||||
class AppSiteInfoResponse(ResponseModel):
|
||||
|
||||
class WebModelConfigResponse(ResponseModel):
|
||||
opening_statement: str | None = None
|
||||
suggested_questions: Any = Field(
|
||||
default=None,
|
||||
validation_alias=AliasChoices("suggested_questions_list", "suggested_questions"),
|
||||
)
|
||||
suggested_questions_after_answer: Any = Field(
|
||||
default=None,
|
||||
validation_alias=AliasChoices("suggested_questions_after_answer_dict", "suggested_questions_after_answer"),
|
||||
)
|
||||
more_like_this: Any = Field(
|
||||
default=None,
|
||||
validation_alias=AliasChoices("more_like_this_dict", "more_like_this"),
|
||||
)
|
||||
model: Any = Field(default=None, validation_alias=AliasChoices("model_dict", "model"))
|
||||
user_input_form: Any = Field(
|
||||
default=None,
|
||||
validation_alias=AliasChoices("user_input_form_list", "user_input_form"),
|
||||
)
|
||||
pre_prompt: str | None = None
|
||||
|
||||
|
||||
class WebAppCustomConfigResponse(ResponseModel):
|
||||
remove_webapp_brand: bool
|
||||
replace_webapp_logo: str | None = None
|
||||
|
||||
|
||||
class WebAppSiteResponse(ResponseModel):
|
||||
app_id: str
|
||||
end_user_id: str | None = None
|
||||
enable_site: bool
|
||||
site: AppSiteResponse
|
||||
model_config_: AppSiteModelConfigResponse | None = Field(default=None, alias="model_config")
|
||||
plan: str | None = None
|
||||
site: WebSiteResponse
|
||||
model_config_: WebModelConfigResponse | None = Field(
|
||||
default=None, validation_alias="model_config", serialization_alias="model_config"
|
||||
)
|
||||
plan: str
|
||||
can_replace_logo: bool
|
||||
custom_config: dict[str, Any] | None = Field(default=None)
|
||||
custom_config: WebAppCustomConfigResponse | None = None
|
||||
|
||||
@classmethod
|
||||
def from_app_site(
|
||||
cls,
|
||||
*,
|
||||
tenant: Tenant,
|
||||
app_model: App,
|
||||
site: Site,
|
||||
end_user_id: str | None,
|
||||
features: FeatureModel,
|
||||
can_replace_logo: bool,
|
||||
) -> Self:
|
||||
custom_config = None
|
||||
if can_replace_logo:
|
||||
replace_webapp_logo = (
|
||||
f"{dify_config.FILES_URL}/files/workspaces/{tenant.id}/webapp-logo"
|
||||
if tenant.custom_config_dict.get("replace_webapp_logo")
|
||||
else None
|
||||
)
|
||||
custom_config = WebAppCustomConfigResponse(
|
||||
remove_webapp_brand=tenant.custom_config_dict.get("remove_webapp_brand", False),
|
||||
replace_webapp_logo=replace_webapp_logo,
|
||||
)
|
||||
|
||||
site_response = WebSiteResponse.model_validate(site, from_attributes=True)
|
||||
if features.billing.enabled and not features.webapp_copyright_enabled:
|
||||
site_response.copyright = None
|
||||
site_response.input_placeholder = None
|
||||
|
||||
return cls(
|
||||
app_id=app_model.id,
|
||||
end_user_id=end_user_id,
|
||||
enable_site=app_model.enable_site,
|
||||
site=site_response,
|
||||
model_config_=None,
|
||||
plan=tenant.plan,
|
||||
can_replace_logo=can_replace_logo,
|
||||
custom_config=custom_config,
|
||||
)
|
||||
|
||||
|
||||
register_response_schema_models(web_ns, AppSiteInfoResponse)
|
||||
register_response_schema_models(
|
||||
web_ns, WebSiteResponse, WebModelConfigResponse, WebAppCustomConfigResponse, WebAppSiteResponse
|
||||
)
|
||||
|
||||
|
||||
@web_ns.route("/site")
|
||||
class AppSiteApi(WebApiResource):
|
||||
"""Resource for app sites."""
|
||||
|
||||
model_config_fields = {
|
||||
"opening_statement": fields.String,
|
||||
"suggested_questions": fields.Raw(attribute="suggested_questions_list"),
|
||||
"suggested_questions_after_answer": fields.Raw(attribute="suggested_questions_after_answer_dict"),
|
||||
"more_like_this": fields.Raw(attribute="more_like_this_dict"),
|
||||
"model": fields.Raw(attribute="model_dict"),
|
||||
"user_input_form": fields.Raw(attribute="user_input_form_list"),
|
||||
"pre_prompt": fields.String,
|
||||
}
|
||||
|
||||
site_fields = {
|
||||
"title": fields.String,
|
||||
"chat_color_theme": fields.String,
|
||||
"chat_color_theme_inverted": fields.Boolean,
|
||||
"icon_type": fields.String,
|
||||
"icon": fields.String,
|
||||
"icon_background": fields.String,
|
||||
"icon_url": AppIconUrlField,
|
||||
"description": fields.String,
|
||||
"copyright": fields.String,
|
||||
"privacy_policy": fields.String,
|
||||
"input_placeholder": fields.String,
|
||||
"custom_disclaimer": fields.String,
|
||||
"default_language": fields.String,
|
||||
"prompt_public": fields.Boolean,
|
||||
"show_workflow_steps": fields.Boolean,
|
||||
"use_icon_as_answer_icon": fields.Boolean,
|
||||
}
|
||||
|
||||
app_fields = {
|
||||
"app_id": fields.String,
|
||||
"end_user_id": fields.String,
|
||||
"enable_site": fields.Boolean,
|
||||
"site": fields.Nested(site_fields),
|
||||
"model_config": fields.Nested(model_config_fields, allow_null=True),
|
||||
"plan": fields.String,
|
||||
"can_replace_logo": fields.Boolean,
|
||||
"custom_config": fields.Raw(attribute="custom_config"),
|
||||
}
|
||||
|
||||
@web_ns.doc("Get App Site Info")
|
||||
@web_ns.doc(description="Retrieve app site information and configuration.")
|
||||
@web_ns.doc(
|
||||
@@ -116,79 +137,26 @@ class AppSiteApi(WebApiResource):
|
||||
500: "Internal Server Error",
|
||||
}
|
||||
)
|
||||
@web_ns.response(200, "Success", web_ns.models[AppSiteInfoResponse.__name__])
|
||||
@marshal_with(app_fields)
|
||||
@web_ns.response(200, "Success", web_ns.models[WebAppSiteResponse.__name__])
|
||||
def get(self, app_model: App, end_user: EndUser):
|
||||
"""Retrieve app site info."""
|
||||
# get site
|
||||
site = db.session.scalar(select(Site).where(Site.app_id == app_model.id).limit(1))
|
||||
|
||||
if not site:
|
||||
if site is None:
|
||||
raise Forbidden()
|
||||
|
||||
if app_model.tenant and app_model.tenant.status == TenantStatus.ARCHIVE:
|
||||
tenant = app_model.tenant
|
||||
if tenant is None or tenant.status == TenantStatus.ARCHIVE:
|
||||
raise Forbidden()
|
||||
|
||||
features = FeatureService.get_features(app_model.tenant_id, exclude_vector_space=True)
|
||||
|
||||
return AppSiteInfo(
|
||||
app_model.tenant,
|
||||
app_model,
|
||||
serialize_runtime_site(site, features),
|
||||
end_user.id,
|
||||
features.can_replace_logo,
|
||||
)
|
||||
|
||||
|
||||
class AppSiteInfo:
|
||||
"""Class to store site information."""
|
||||
|
||||
def __init__(self, tenant, app, site, end_user, can_replace_logo):
|
||||
"""Initialize AppSiteInfo instance."""
|
||||
self.app_id = app.id
|
||||
self.end_user_id = end_user
|
||||
self.enable_site = app.enable_site
|
||||
self.site = site
|
||||
self.model_config = None
|
||||
self.plan = tenant.plan
|
||||
self.can_replace_logo = can_replace_logo
|
||||
|
||||
if can_replace_logo:
|
||||
base_url = dify_config.FILES_URL
|
||||
remove_webapp_brand = tenant.custom_config_dict.get("remove_webapp_brand", False)
|
||||
replace_webapp_logo = (
|
||||
f"{base_url}/files/workspaces/{tenant.id}/webapp-logo"
|
||||
if tenant.custom_config_dict.get("replace_webapp_logo")
|
||||
else None
|
||||
)
|
||||
self.custom_config = {
|
||||
"remove_webapp_brand": remove_webapp_brand,
|
||||
"replace_webapp_logo": replace_webapp_logo,
|
||||
}
|
||||
|
||||
|
||||
def serialize_site(site: Site) -> dict[str, Any]:
|
||||
"""Serialize Site model using the same schema as AppSiteApi."""
|
||||
return cast(dict[str, Any], marshal(site, AppSiteApi.site_fields))
|
||||
|
||||
|
||||
def serialize_runtime_site(site: Site, features: FeatureModel) -> dict[str, Any]:
|
||||
site_payload = serialize_site(site)
|
||||
if not features.billing.enabled or features.webapp_copyright_enabled:
|
||||
return site_payload
|
||||
|
||||
site_payload["copyright"] = None
|
||||
site_payload["input_placeholder"] = None
|
||||
return site_payload
|
||||
|
||||
|
||||
def serialize_app_site_payload(app_model: App, site: Site, end_user_id: str | None) -> dict[str, Any]:
|
||||
features = FeatureService.get_features(app_model.tenant_id, exclude_vector_space=True)
|
||||
app_site_info = AppSiteInfo(
|
||||
app_model.tenant,
|
||||
app_model,
|
||||
serialize_runtime_site(site, features),
|
||||
end_user_id,
|
||||
features.can_replace_logo,
|
||||
)
|
||||
return cast(dict[str, Any], marshal(app_site_info, AppSiteApi.app_fields))
|
||||
return WebAppSiteResponse.from_app_site(
|
||||
tenant=tenant,
|
||||
app_model=app_model,
|
||||
site=site,
|
||||
end_user_id=end_user.id,
|
||||
features=features,
|
||||
can_replace_logo=features.can_replace_logo,
|
||||
).model_dump(mode="json")
|
||||
|
||||
@@ -76,6 +76,7 @@ class WorkflowRunApi(WebApiResource):
|
||||
streaming=True,
|
||||
)
|
||||
|
||||
# response-contract:ignore compact_generate_response
|
||||
return helper.compact_generate_response(response)
|
||||
except ProviderTokenNotInitError as ex:
|
||||
raise ProviderNotInitializeError(ex.description)
|
||||
@@ -129,4 +130,4 @@ class WorkflowTaskStopApi(WebApiResource):
|
||||
# New graph engine command channel mechanism
|
||||
GraphEngineManager(redis_client).send_stop_command(task_id)
|
||||
|
||||
return {"result": "success"}
|
||||
return SimpleResultResponse(result="success").model_dump(mode="json")
|
||||
|
||||
Reference in New Issue
Block a user