Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
wraps.py241 linesDownload Raw Back to service_api
1from collections.abc import Callable2from datetime import datetime, timezone3from enum import Enum4from functools import wraps5from typing import Optional6 7from flask import current_app, request8from flask_login import user_logged_in9from flask_restful import Resource10from pydantic import BaseModel11from werkzeug.exceptions import Forbidden, Unauthorized12 13from extensions.ext_database import db14from libs.login import _get_user15from models.account import Account, Tenant, TenantAccountJoin, TenantStatus16from models.model import ApiToken, App, EndUser17from services.feature_service import FeatureService18 19 20class WhereisUserArg(Enum):21    """22    Enum for whereis_user_arg.23    """24 25    QUERY = "query"26    JSON = "json"27    FORM = "form"28 29 30class FetchUserArg(BaseModel):31    fetch_from: WhereisUserArg32    required: bool = False33 34 35def validate_app_token(view: Optional[Callable] = None, *, fetch_user_arg: Optional[FetchUserArg] = None):36    def decorator(view_func):37        @wraps(view_func)38        def decorated_view(*args, **kwargs):39            api_token = validate_and_get_api_token("app")40 41            app_model = db.session.query(App).filter(App.id == api_token.app_id).first()42            if not app_model:43                raise Forbidden("The app no longer exists.")44 45            if app_model.status != "normal":46                raise Forbidden("The app's status is abnormal.")47 48            if not app_model.enable_api:49                raise Forbidden("The app's API service has been disabled.")50 51            tenant = db.session.query(Tenant).filter(Tenant.id == app_model.tenant_id).first()52            if tenant.status == TenantStatus.ARCHIVE:53                raise Forbidden("The workspace's status is archived.")54 55            kwargs["app_model"] = app_model56 57            if fetch_user_arg:58                if fetch_user_arg.fetch_from == WhereisUserArg.QUERY:59                    user_id = request.args.get("user")60                elif fetch_user_arg.fetch_from == WhereisUserArg.JSON:61                    user_id = request.get_json().get("user")62                elif fetch_user_arg.fetch_from == WhereisUserArg.FORM:63                    user_id = request.form.get("user")64                else:65                    # use default-user66                    user_id = None67 68                if not user_id and fetch_user_arg.required:69                    raise ValueError("Arg user must be provided.")70 71                if user_id:72                    user_id = str(user_id)73 74                kwargs["end_user"] = create_or_update_end_user_for_user_id(app_model, user_id)75 76            return view_func(*args, **kwargs)77 78        return decorated_view79 80    if view is None:81        return decorator82    else:83        return decorator(view)84 85 86def cloud_edition_billing_resource_check(resource: str, api_token_type: str):87    def interceptor(view):88        def decorated(*args, **kwargs):89            api_token = validate_and_get_api_token(api_token_type)90            features = FeatureService.get_features(api_token.tenant_id)91 92            if features.billing.enabled:93                members = features.members94                apps = features.apps95                vector_space = features.vector_space96                documents_upload_quota = features.documents_upload_quota97 98                if resource == "members" and 0 < members.limit <= members.size:99                    raise Forbidden("The number of members has reached the limit of your subscription.")100                elif resource == "apps" and 0 < apps.limit <= apps.size:101                    raise Forbidden("The number of apps has reached the limit of your subscription.")102                elif resource == "vector_space" and 0 < vector_space.limit <= vector_space.size:103                    raise Forbidden("The capacity of the vector space has reached the limit of your subscription.")104                elif resource == "documents" and 0 < documents_upload_quota.limit <= documents_upload_quota.size:105                    raise Forbidden("The number of documents has reached the limit of your subscription.")106                else:107                    return view(*args, **kwargs)108 109            return view(*args, **kwargs)110 111        return decorated112 113    return interceptor114 115 116def cloud_edition_billing_knowledge_limit_check(resource: str, api_token_type: str):117    def interceptor(view):118        @wraps(view)119        def decorated(*args, **kwargs):120            api_token = validate_and_get_api_token(api_token_type)121            features = FeatureService.get_features(api_token.tenant_id)122            if features.billing.enabled:123                if resource == "add_segment":124                    if features.billing.subscription.plan == "sandbox":125                        raise Forbidden(126                            "To unlock this feature and elevate your Dify experience, please upgrade to a paid plan."127                        )128                else:129                    return view(*args, **kwargs)130 131            return view(*args, **kwargs)132 133        return decorated134 135    return interceptor136 137 138def validate_dataset_token(view=None):139    def decorator(view):140        @wraps(view)141        def decorated(*args, **kwargs):142            api_token = validate_and_get_api_token("dataset")143            tenant_account_join = (144                db.session.query(Tenant, TenantAccountJoin)145                .filter(Tenant.id == api_token.tenant_id)146                .filter(TenantAccountJoin.tenant_id == Tenant.id)147                .filter(TenantAccountJoin.role.in_(["owner"]))148                .filter(Tenant.status == TenantStatus.NORMAL)149                .one_or_none()150            )  # TODO: only owner information is required, so only one is returned.151            if tenant_account_join:152                tenant, ta = tenant_account_join153                account = Account.query.filter_by(id=ta.account_id).first()154                # Login admin155                if account:156                    account.current_tenant = tenant157                    current_app.login_manager._update_request_context_with_user(account)158                    user_logged_in.send(current_app._get_current_object(), user=_get_user())159                else:160                    raise Unauthorized("Tenant owner account does not exist.")161            else:162                raise Unauthorized("Tenant does not exist.")163            return view(api_token.tenant_id, *args, **kwargs)164 165        return decorated166 167    if view:168        return decorator(view)169 170    # if view is None, it means that the decorator is used without parentheses171    # use the decorator as a function for method_decorators172    return decorator173 174 175def validate_and_get_api_token(scope=None):176    """177    Validate and get API token.178    """179    auth_header = request.headers.get("Authorization")180    if auth_header is None or " " not in auth_header:181        raise Unauthorized("Authorization header must be provided and start with 'Bearer'")182 183    auth_scheme, auth_token = auth_header.split(None, 1)184    auth_scheme = auth_scheme.lower()185 186    if auth_scheme != "bearer":187        raise Unauthorized("Authorization scheme must be 'Bearer'")188 189    api_token = (190        db.session.query(ApiToken)191        .filter(192            ApiToken.token == auth_token,193            ApiToken.type == scope,194        )195        .first()196    )197 198    if not api_token:199        raise Unauthorized("Access token is invalid")200 201    api_token.last_used_at = datetime.now(timezone.utc).replace(tzinfo=None)202    db.session.commit()203 204    return api_token205 206 207def create_or_update_end_user_for_user_id(app_model: App, user_id: Optional[str] = None) -> EndUser:208    """209    Create or update session terminal based on user ID.210    """211    if not user_id:212        user_id = "DEFAULT-USER"213 214    end_user = (215        db.session.query(EndUser)216        .filter(217            EndUser.tenant_id == app_model.tenant_id,218            EndUser.app_id == app_model.id,219            EndUser.session_id == user_id,220            EndUser.type == "service_api",221        )222        .first()223    )224 225    if end_user is None:226        end_user = EndUser(227            tenant_id=app_model.tenant_id,228            app_id=app_model.id,229            type="service_api",230            is_anonymous=True if user_id == "DEFAULT-USER" else False,231            session_id=user_id,232        )233        db.session.add(end_user)234        db.session.commit()235 236    return end_user237 238 239class DatasetApiResource(Resource):240    method_decorators = [validate_dataset_token]241