mirror of
https://github.com/langgenius/dify.git
synced 2026-09-24 23:22:26 +08:00
refactor: accept db.session explicitly in APIBasedExtensionService (#37693)
This commit is contained in:
@@ -1,16 +1,16 @@
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from core.extension.api_based_extension_requestor import APIBasedExtensionRequestor
|
||||
from core.helper.encrypter import decrypt_token, encrypt_token
|
||||
from extensions.ext_database import db
|
||||
from models.api_based_extension import APIBasedExtension, APIBasedExtensionPoint
|
||||
|
||||
|
||||
class APIBasedExtensionService:
|
||||
@staticmethod
|
||||
def get_all_by_tenant_id(tenant_id: str) -> list[APIBasedExtension]:
|
||||
def get_all_by_tenant_id(session: Session, tenant_id: str) -> list[APIBasedExtension]:
|
||||
extension_list = list(
|
||||
db.session.scalars(
|
||||
session.scalars(
|
||||
select(APIBasedExtension)
|
||||
.where(APIBasedExtension.tenant_id == tenant_id)
|
||||
.order_by(APIBasedExtension.created_at.desc())
|
||||
@@ -23,23 +23,23 @@ class APIBasedExtensionService:
|
||||
return extension_list
|
||||
|
||||
@classmethod
|
||||
def save(cls, extension_data: APIBasedExtension) -> APIBasedExtension:
|
||||
cls._validation(extension_data)
|
||||
def save(cls, session: Session, extension_data: APIBasedExtension) -> APIBasedExtension:
|
||||
cls._validation(session, extension_data)
|
||||
|
||||
extension_data.api_key = encrypt_token(extension_data.tenant_id, extension_data.api_key)
|
||||
|
||||
db.session.add(extension_data)
|
||||
db.session.commit()
|
||||
session.add(extension_data)
|
||||
session.commit()
|
||||
return extension_data
|
||||
|
||||
@staticmethod
|
||||
def delete(extension_data: APIBasedExtension):
|
||||
db.session.delete(extension_data)
|
||||
db.session.commit()
|
||||
def delete(session: Session, extension_data: APIBasedExtension):
|
||||
session.delete(extension_data)
|
||||
session.commit()
|
||||
|
||||
@staticmethod
|
||||
def get_with_tenant_id(tenant_id: str, api_based_extension_id: str) -> APIBasedExtension:
|
||||
extension = db.session.scalar(
|
||||
def get_with_tenant_id(session: Session, tenant_id: str, api_based_extension_id: str) -> APIBasedExtension:
|
||||
extension = session.scalar(
|
||||
select(APIBasedExtension)
|
||||
.where(APIBasedExtension.tenant_id == tenant_id, APIBasedExtension.id == api_based_extension_id)
|
||||
.limit(1)
|
||||
@@ -53,14 +53,14 @@ class APIBasedExtensionService:
|
||||
return extension
|
||||
|
||||
@classmethod
|
||||
def _validation(cls, extension_data: APIBasedExtension):
|
||||
def _validation(cls, session: Session, extension_data: APIBasedExtension):
|
||||
# name
|
||||
if not extension_data.name:
|
||||
raise ValueError("name must not be empty")
|
||||
|
||||
if not extension_data.id:
|
||||
# case one: check new data, name must be unique
|
||||
is_name_existed = db.session.scalar(
|
||||
is_name_existed = session.scalar(
|
||||
select(APIBasedExtension)
|
||||
.where(
|
||||
APIBasedExtension.tenant_id == extension_data.tenant_id,
|
||||
@@ -73,7 +73,7 @@ class APIBasedExtensionService:
|
||||
raise ValueError("name must be unique, it is already existed")
|
||||
else:
|
||||
# case two: check existing data, name must be unique
|
||||
is_name_existed = db.session.scalar(
|
||||
is_name_existed = session.scalar(
|
||||
select(APIBasedExtension)
|
||||
.where(
|
||||
APIBasedExtension.tenant_id == extension_data.tenant_id,
|
||||
|
||||
Reference in New Issue
Block a user