refactor: manage rag pipeline sessions explicitly (#38274)

This commit is contained in:
Byron.wang
2026-07-03 08:15:37 +00:00
committed by GitHub
parent 262b0b1a89
commit 5cb76f5eff
25 changed files with 972 additions and 793 deletions
@@ -5,7 +5,7 @@ from flask import request
from flask_restx import Resource
from pydantic import BaseModel, Field
from sqlalchemy import select
from sqlalchemy.orm import sessionmaker
from sqlalchemy.orm import Session, sessionmaker
from werkzeug.exceptions import NotFound
from controllers.common.fields import SimpleDataResponse
@@ -16,6 +16,7 @@ from controllers.common.schema import (
register_schema_models,
)
from controllers.console import console_ns
from controllers.console.app.wraps import with_session
from controllers.console.wraps import (
account_initialization_required,
enterprise_license_required,
@@ -102,10 +103,13 @@ class PipelineTemplateListApi(Resource):
@account_initialization_required
@enterprise_license_required
@with_current_tenant_id
def get(self, current_tenant_id: str) -> JsonResponseWithStatus:
@with_session
def get(self, session: Session, current_tenant_id: str) -> JsonResponseWithStatus:
query = PipelineTemplateListQuery.model_validate(request.args.to_dict(flat=True))
# get pipeline templates
pipeline_templates = RagPipelineService.get_pipeline_templates(query.type, query.language, current_tenant_id)
pipeline_templates = RagPipelineService.get_pipeline_templates(
session, query.type, query.language, current_tenant_id
)
return dump_response(PipelineTemplateListResponse, pipeline_templates), 200
@@ -117,10 +121,11 @@ class PipelineTemplateDetailApi(Resource):
@login_required
@account_initialization_required
@enterprise_license_required
def get(self, template_id: str) -> JsonResponseWithStatus:
@with_session
def get(self, session: Session, template_id: str) -> JsonResponseWithStatus:
query = PipelineTemplateDetailQuery.model_validate(request.args.to_dict(flat=True))
rag_pipeline_service = RagPipelineService()
pipeline_template = rag_pipeline_service.get_pipeline_template_detail(template_id, query.type)
pipeline_template = rag_pipeline_service.get_pipeline_template_detail(session, template_id, query.type)
if pipeline_template is None:
raise NotFound("Pipeline template not found from upstream service.")
return dump_response(PipelineTemplateDetailResponse, pipeline_template), 200
@@ -6,7 +6,7 @@ from uuid import UUID
from flask import abort, request
from flask_restx import Resource
from pydantic import BaseModel, Field, RootModel, ValidationError
from sqlalchemy.orm import sessionmaker
from sqlalchemy.orm import Session, sessionmaker
from werkzeug.exceptions import BadRequest, Forbidden, InternalServerError, NotFound
import services
@@ -26,6 +26,7 @@ from controllers.console.app.workflow import (
WorkflowPaginationResponse,
WorkflowResponse,
)
from controllers.console.app.wraps import with_session
from controllers.console.datasets.wraps import get_rag_pipeline
from controllers.console.wraps import (
RBACPermission,
@@ -343,7 +344,8 @@ class DraftRagPipelineRunApi(Resource):
@rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT)
@with_current_user
@get_rag_pipeline
def post(self, current_user: Account, pipeline: Pipeline):
@with_session
def post(self, session: Session, current_user: Account, pipeline: Pipeline):
"""
Run draft workflow
"""
@@ -352,6 +354,7 @@ class DraftRagPipelineRunApi(Resource):
try:
response = PipelineGenerateService.generate(
session=session,
pipeline=pipeline,
user=current_user,
args=args,
@@ -375,7 +378,8 @@ class PublishedRagPipelineRunApi(Resource):
@rbac_permission_required(RBACResourceScope.DATASET, RBACPermission.DATASET_EDIT)
@with_current_user
@get_rag_pipeline
def post(self, current_user: Account, pipeline: Pipeline):
@with_session
def post(self, session: Session, current_user: Account, pipeline: Pipeline):
"""
Run published workflow
"""
@@ -385,6 +389,7 @@ class PublishedRagPipelineRunApi(Resource):
try:
response = PipelineGenerateService.generate(
session=session,
pipeline=pipeline,
user=current_user,
args=args,
@@ -1014,13 +1019,14 @@ class RagPipelineTransformApi(Resource):
@login_required
@account_initialization_required
@with_current_user
def post(self, current_user: Account, dataset_id: UUID):
@with_session
def post(self, session: Session, current_user: Account, dataset_id: UUID):
if not (current_user.has_edit_permission or current_user.is_dataset_operator):
raise Forbidden()
dataset_id_str = str(dataset_id)
rag_pipeline_transform_service = RagPipelineTransformService()
result = rag_pipeline_transform_service.transform_dataset(dataset_id_str, db.session)
result = rag_pipeline_transform_service.transform_dataset(dataset_id_str, session)
return result
@@ -5,6 +5,7 @@ from uuid import UUID
from flask import request
from pydantic import BaseModel, Field, RootModel
from sqlalchemy import select
from sqlalchemy.orm import Session
from werkzeug.exceptions import Forbidden, NotFound
import services
@@ -16,6 +17,7 @@ from controllers.common.schema import (
register_schema_model,
register_schema_models,
)
from controllers.console.app.wraps import with_session
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
@@ -264,7 +266,8 @@ class PipelineRunApi(DatasetApiResource):
"Pipeline run successfully",
service_api_ns.models[GeneratedAppResponse.__name__],
)
def post(self, tenant_id: str, dataset_id: UUID):
@with_session
def post(self, session: Session, tenant_id: str, dataset_id: UUID):
"""Resource for running a rag pipeline."""
dataset_id_str = str(dataset_id)
# Verify dataset ownership
@@ -282,6 +285,7 @@ class PipelineRunApi(DatasetApiResource):
pipeline: Pipeline = rag_pipeline_service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset_id_str)
try:
response: dict[Any, Any] | Generator[str, Any, None] = PipelineGenerateService.generate(
session=session,
pipeline=pipeline,
user=current_user,
args=payload.model_dump(),