refactor: accept db.session explicitly in APIBasedExtensionService (#37693)

This commit is contained in:
Rohit Gahlawat
2026-06-21 00:53:36 +00:00
committed by GitHub
parent 75d50455d6
commit 9b4dd9d4e8
5 changed files with 85 additions and 65 deletions
+15 -15
View File
@@ -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,