refactor: explicit DB session propagation across backend paths (#38559)

Co-authored-by: WH-2099 <wh2099@pm.me>
This commit is contained in:
Byron.wang
2026-07-15 06:48:28 +00:00
committed by GitHub
co-authored by WH-2099
parent 7d5835fbc0
commit ab3e4daa95
397 changed files with 14205 additions and 9967 deletions
+19 -13
View File
@@ -42,7 +42,7 @@ from graphon.model_runtime.entities.message_entities import ImagePromptMessageCo
from graphon.model_runtime.entities.model_entities import ModelFeature
from graphon.model_runtime.model_providers.base.large_language_model import LargeLanguageModel
from models.enums import CreatorUserRole
from models.model import Conversation, Message, MessageAgentThought, MessageFile
from models.model import Conversation, Message, MessageAgentThought, MessageFile, load_annotation_reply_config
logger = logging.getLogger(__name__)
_file_access_controller = DatabaseFileAccessController()
@@ -76,7 +76,9 @@ class BaseAgentRunner(AppRunner):
self.message = message
self.user_id = user_id
self.memory = memory
self.history_prompt_messages = self.organize_agent_history(prompt_messages=prompt_messages or [])
self.history_prompt_messages = self.organize_agent_history(
session=session, prompt_messages=prompt_messages or []
)
self.model_instance = model_instance
# init callback
@@ -104,7 +106,7 @@ class BaseAgentRunner(AppRunner):
)
# get how many agent thoughts have been created
self.agent_thought_count = (
db.session.scalar(
session.scalar(
select(func.count())
.select_from(MessageAgentThought)
.where(
@@ -113,7 +115,7 @@ class BaseAgentRunner(AppRunner):
)
or 0
)
db.session.close()
session.close()
# check if model supports stream tool call
llm_model = cast(LargeLanguageModel, model_instance.model_type_instance)
@@ -350,7 +352,7 @@ class BaseAgentRunner(AppRunner):
db.session.commit()
db.session.close()
def organize_agent_history(self, prompt_messages: list[PromptMessage]) -> list[PromptMessage]:
def organize_agent_history(self, prompt_messages: list[PromptMessage], *, session: Session) -> list[PromptMessage]:
"""
Organize agent history
"""
@@ -362,7 +364,7 @@ class BaseAgentRunner(AppRunner):
messages = (
(
db.session.execute(
session.execute(
select(Message)
.where(Message.conversation_id == self.message.conversation_id)
.order_by(Message.created_at.desc())
@@ -378,8 +380,8 @@ class BaseAgentRunner(AppRunner):
if message.id == self.message.id:
continue
result.append(self.organize_agent_user_prompt(message))
agent_thoughts = message.agent_thoughts
result.append(self.organize_agent_user_prompt(message, session=session))
agent_thoughts = message.agent_thoughts_with_session(session=session)
if agent_thoughts:
for agent_thought in agent_thoughts:
tool_names_raw = agent_thought.tool
@@ -441,17 +443,21 @@ class BaseAgentRunner(AppRunner):
if message.answer:
result.append(AssistantPromptMessage(content=message.answer))
db.session.close()
session.close()
return result
def organize_agent_user_prompt(self, message: Message) -> UserPromptMessage:
def organize_agent_user_prompt(self, message: Message, *, session: Session) -> UserPromptMessage:
stmt = select(MessageFile).where(MessageFile.message_id == message.id)
files = db.session.scalars(stmt).all()
files = session.scalars(stmt).all()
if not files:
return UserPromptMessage(content=message.query)
if message.app_model_config:
file_extra_config = FileUploadConfigManager.convert(message.app_model_config.to_dict())
app_model_config = message.app_model_config_with_session(session=session)
if app_model_config:
annotation_reply = load_annotation_reply_config(session, app_model_config.app_id)
file_extra_config = FileUploadConfigManager.convert(
app_model_config.to_dict(annotation_reply=annotation_reply)
)
else:
file_extra_config = None
+12 -1
View File
@@ -114,7 +114,11 @@ class CotAgentRunner(BaseAgentRunner, ABC):
message_file_ids: list[str] = []
agent_thought_id = self.create_agent_thought(
message_id=message.id, message="", tool_name="", tool_input="", messages_ids=message_file_ids
message_id=message.id,
message="",
tool_name="",
tool_input="",
messages_ids=message_file_ids,
)
if iteration_step > 1:
@@ -125,6 +129,11 @@ class CotAgentRunner(BaseAgentRunner, ABC):
# recalc llm max tokens
prompt_messages = self._organize_prompt_messages()
self.recalc_llm_max_tokens(self.model_config, prompt_messages)
# Release any setup/tool transaction before waiting on the provider stream.
session.commit()
session.close()
# invoke model
chunks = model_instance.invoke_llm(
prompt_messages=prompt_messages,
@@ -333,6 +342,8 @@ class CotAgentRunner(BaseAgentRunner, ABC):
agent_tool_callback=self.agent_callback,
trace_manager=trace_manager,
)
session.commit()
session.close()
# publish files
for message_file_id in message_files:
+12 -1
View File
@@ -87,12 +87,21 @@ class FunctionCallAgentRunner(BaseAgentRunner):
message_file_ids: list[str] = []
agent_thought_id = self.create_agent_thought(
message_id=message.id, message="", tool_name="", tool_input="", messages_ids=message_file_ids
message_id=message.id,
message="",
tool_name="",
tool_input="",
messages_ids=message_file_ids,
)
# recalc llm max tokens
prompt_messages = self._organize_prompt_messages()
self.recalc_llm_max_tokens(self.model_config, prompt_messages)
# Release any setup/tool transaction before waiting on the provider stream.
session.commit()
session.close()
# invoke model
chunks: Union[Generator[LLMResultChunk, None, None], LLMResult] = model_instance.invoke_llm(
prompt_messages=prompt_messages,
@@ -256,6 +265,8 @@ class FunctionCallAgentRunner(BaseAgentRunner):
message_id=self.message.id,
conversation_id=self.conversation.id,
)
session.commit()
session.close()
# publish files
for message_file_id in message_files:
# publish message file
@@ -1,6 +1,8 @@
import uuid
from typing import Any, Literal, cast
from sqlalchemy.orm import Session
from core.app.app_config.entities import (
DatasetEntity,
DatasetRetrieveConfigEntity,
@@ -9,7 +11,6 @@ from core.app.app_config.entities import (
)
from core.entities.agent_entities import PlanningStrategy
from core.rag.data_post_processor.data_post_processor import RerankingModelDict, WeightsDict
from extensions.ext_database import db
from models.model import AppMode, AppModelConfigDict
from services.dataset_service import DatasetService
@@ -140,7 +141,7 @@ class DatasetConfigManager:
@classmethod
def validate_and_set_defaults(
cls, tenant_id: str, app_mode: AppMode, config: dict[str, Any]
cls, tenant_id: str, app_mode: AppMode, config: dict[str, Any], session: Session
) -> tuple[dict[str, Any], list[str]]:
"""
Validate and set defaults for dataset feature
@@ -150,7 +151,7 @@ class DatasetConfigManager:
:param config: app model config args
"""
# Extract dataset config for legacy compatibility
config = cls.extract_dataset_config_for_legacy_compatibility(tenant_id, app_mode, config)
config = cls.extract_dataset_config_for_legacy_compatibility(tenant_id, app_mode, config, session)
# dataset_configs
if "dataset_configs" not in config or not config.get("dataset_configs"):
@@ -175,7 +176,9 @@ class DatasetConfigManager:
return config, ["agent_mode", "dataset_configs", "dataset_query_variable"]
@classmethod
def extract_dataset_config_for_legacy_compatibility(cls, tenant_id: str, app_mode: AppMode, config: dict[str, Any]):
def extract_dataset_config_for_legacy_compatibility(
cls, tenant_id: str, app_mode: AppMode, config: dict[str, Any], session: Session
):
"""
Extract dataset config for legacy compatibility
@@ -238,7 +241,7 @@ class DatasetConfigManager:
except ValueError:
raise ValueError("id in dataset must be of UUID type")
if not cls.is_dataset_exists(tenant_id, tool_item["id"]):
if not cls.is_dataset_exists(tenant_id, tool_item["id"], session):
raise ValueError("Dataset ID does not exist, please check your permission.")
has_datasets = True
@@ -255,9 +258,9 @@ class DatasetConfigManager:
return config
@classmethod
def is_dataset_exists(cls, tenant_id: str, dataset_id: str) -> bool:
def is_dataset_exists(cls, tenant_id: str, dataset_id: str, session: Session) -> bool:
# verify if the dataset ID exists
dataset = DatasetService.get_dataset(dataset_id, db.session())
dataset = DatasetService.get_dataset(dataset_id, session)
if not dataset:
return False
@@ -86,6 +86,8 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
workflow_run_id: str,
streaming: Literal[False],
pause_state_config: PauseStateLayerConfig | None = None,
*,
session: Session,
) -> Mapping[str, Any]: ...
@overload
@@ -99,6 +101,8 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
workflow_run_id: str,
streaming: Literal[True],
pause_state_config: PauseStateLayerConfig | None = None,
*,
session: Session,
) -> Generator[Mapping | str, None, None]: ...
@overload
@@ -112,6 +116,8 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
workflow_run_id: str,
streaming: bool,
pause_state_config: PauseStateLayerConfig | None = None,
*,
session: Session,
) -> Mapping[str, Any] | Generator[str | Mapping, None, None]: ...
def generate(
@@ -124,6 +130,8 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
workflow_run_id: str,
streaming: bool = True,
pause_state_config: PauseStateLayerConfig | None = None,
*,
session: Session,
) -> Mapping[str, Any] | Generator[str | Mapping, None, None]:
"""
Generate App response.
@@ -134,6 +142,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
:param args: request args
:param invoke_from: invoke from source
:param streaming: is stream
:param session: database session supplied by the caller
"""
if not args.get("query"):
raise ValueError("query is required")
@@ -157,7 +166,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
if conversation_id:
try:
conversation = ConversationService.get_conversation(
app_model=app_model, conversation_id=conversation_id, user=user, session=db.session()
app_model=app_model, conversation_id=conversation_id, user=user, session=session
)
except ConversationNotExistsError:
if invoke_from == InvokeFrom.SERVICE_API:
@@ -255,6 +264,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
conversation=conversation,
stream=streaming,
pause_state_config=pause_state_config,
session=session,
)
def resume(
@@ -265,6 +275,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
user: Account | EndUser,
conversation: Conversation,
message: Message,
session: Session,
application_generate_entity: AdvancedChatAppGenerateEntity,
workflow_execution_repository: WorkflowExecutionRepository,
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
@@ -301,6 +312,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
pause_state_config=pause_state_config,
graph_runtime_state=graph_runtime_state,
response_stream_filter=response_stream_filter,
session=session,
)
def single_iteration_generate(
@@ -311,6 +323,8 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
user: Account | EndUser,
args: Mapping[str, Any],
streaming: bool = True,
*,
session: Session,
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]:
"""
Generate App response.
@@ -321,6 +335,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
:param user: account or end user
:param args: request args
:param streaming: is streamed
:param session: database session supplied by the caller
"""
if not node_id:
raise ValueError("node_id is required")
@@ -377,7 +392,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
tenant_id=application_generate_entity.app_config.tenant_id,
user_id=user.id,
)
draft_var_srv = WorkflowDraftVariableService(db.session())
draft_var_srv = WorkflowDraftVariableService(session)
draft_var_srv.prefill_conversation_variable_default_values(workflow, user_id=user.id)
return self._generate(
@@ -390,6 +405,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
conversation=None,
stream=streaming,
variable_loader=var_loader,
session=session,
)
def single_loop_generate(
@@ -400,6 +416,8 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
user: Account | EndUser,
args: LoopNodeRunPayload,
streaming: bool = True,
*,
session: Session,
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]:
"""
Generate App response.
@@ -410,6 +428,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
:param user: account or end user
:param args: request args
:param streaming: is stream
:param session: database session supplied by the caller
"""
if not node_id:
raise ValueError("node_id is required")
@@ -464,7 +483,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
tenant_id=application_generate_entity.app_config.tenant_id,
user_id=user.id,
)
draft_var_srv = WorkflowDraftVariableService(db.session())
draft_var_srv = WorkflowDraftVariableService(session)
draft_var_srv.prefill_conversation_variable_default_values(workflow, user_id=user.id)
return self._generate(
@@ -477,6 +496,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
conversation=None,
stream=streaming,
variable_loader=var_loader,
session=session,
)
def _generate(
@@ -486,6 +506,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
user: Account | EndUser,
invoke_from: InvokeFrom,
application_generate_entity: AdvancedChatAppGenerateEntity,
session: Session,
workflow_execution_repository: WorkflowExecutionRepository,
workflow_node_execution_repository: WorkflowNodeExecutionRepository,
conversation: Conversation | None = None,
@@ -504,6 +525,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
:param user: account or end user
:param invoke_from: invoke from source
:param application_generate_entity: application generate entity
:param session: database session supplied by the caller
:param workflow_execution_repository: repository for workflow execution
:param workflow_node_execution_repository: repository for workflow node execution
:param conversation: conversation
@@ -519,18 +541,22 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
if conversation is not None and message is not None:
pass
else:
conversation, message = self._init_generate_records(application_generate_entity, conversation)
conversation, message = self._init_generate_records(
application_generate_entity,
conversation,
session=session,
)
if is_first_conversation:
# update conversation features
conversation.override_model_configs = workflow.features
db.session.commit()
db.session.refresh(conversation)
session.commit()
session.refresh(conversation)
# get conversation dialogue count
# NOTE: dialogue_count should not start from 0,
# because during the first conversation, dialogue_count should be 1.
self._dialogue_count = get_thread_messages_length(conversation.id) + 1
self._dialogue_count = get_thread_messages_length(conversation.id, session=session) + 1
# init queue manager
queue_manager = MessageBasedAppQueueManager(
@@ -582,7 +608,7 @@ class AdvancedChatAppGenerator(MessageBasedAppGenerator):
workflow_snapshot = WorkflowSnapshot.from_workflow(workflow)
conversation_snapshot = ConversationSnapshot.from_conversation(conversation)
message_snapshot = MessageSnapshot.from_message(message)
db.session.close()
session.close()
# return response or stream generator
response = self._handle_advanced_chat_response(
+39 -28
View File
@@ -39,7 +39,6 @@ from core.workflow.system_variables import (
)
from core.workflow.variable_pool_initializer import add_node_inputs_to_pool, add_variables_to_pool
from core.workflow.workflow_entry import WorkflowEntry
from extensions.ext_database import db
from extensions.ext_redis import redis_client
from extensions.otel import WorkflowAppRunnerHandler, trace_span
from extensions.workflow_warm_shutdown import WORKFLOW_WARM_SHUTDOWN_ABORT_REASON, celery_warm_shutdown_started
@@ -173,12 +172,21 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
)
# annotation reply
if self.handle_annotation_reply(
app_record=self._app,
message=self.message,
query=new_query,
app_generate_entity=self.application_generate_entity,
):
with create_session() as session:
annotation_reply = self.handle_annotation_reply(
app_record=self._app,
message=self.message,
query=new_query,
app_generate_entity=self.application_generate_entity,
session=session,
)
session.commit()
if annotation_reply:
self._publish_event(QueueAnnotationReplyEvent(message_annotation_id=annotation_reply.id))
self._complete_with_stream_output(
text=annotation_reply.content,
stopped_by=QueueStopEvent.StopBy.ANNOTATION_REPLY,
)
return
# Initialize conversation variables
@@ -212,10 +220,6 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
trace_session_id=self.application_generate_entity.extras.get("trace_session_id"),
)
# Release the Flask scoped session before workflow execution so a checked-out DB connection
# is not held for the lifetime of the graph run.
db.session.close()
# RUN WORKFLOW
# Create Redis command channel for this workflow execution
task_id = self.application_generate_entity.task_id
@@ -300,26 +304,22 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
return False, new_inputs, new_query
def handle_annotation_reply(
self, app_record: App, message: Message, query: str, app_generate_entity: AdvancedChatAppGenerateEntity
) -> bool:
annotation_reply = self.query_app_annotations_to_reply(
self,
app_record: App,
message: Message,
query: str,
app_generate_entity: AdvancedChatAppGenerateEntity,
session: Session,
) -> MessageAnnotation | None:
return self.query_app_annotations_to_reply(
app_record=app_record,
message=message,
query=query,
user_id=app_generate_entity.user_id,
invoke_from=app_generate_entity.invoke_from,
session=session,
)
if annotation_reply:
self._publish_event(QueueAnnotationReplyEvent(message_annotation_id=annotation_reply.id))
self._complete_with_stream_output(
text=annotation_reply.content, stopped_by=QueueStopEvent.StopBy.ANNOTATION_REPLY
)
return True
return False
def _complete_with_stream_output(self, text: str, stopped_by: QueueStopEvent.StopBy):
"""
Direct output
@@ -329,7 +329,13 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
self._publish_event(QueueStopEvent(stopped_by=stopped_by))
def query_app_annotations_to_reply(
self, app_record: App, message: Message, query: str, user_id: str, invoke_from: InvokeFrom
self,
app_record: App,
message: Message,
query: str,
user_id: str,
invoke_from: InvokeFrom,
session: Session,
) -> MessageAnnotation | None:
"""
Query app annotations to reply
@@ -342,7 +348,12 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
"""
annotation_reply_feature = AnnotationReplyFeature()
return annotation_reply_feature.query(
app_record=app_record, message=message, query=query, user_id=user_id, invoke_from=invoke_from
app_record=app_record,
message=message,
query=query,
user_id=user_id,
invoke_from=invoke_from,
session=session,
)
def moderation_for_inputs(
@@ -395,7 +406,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
existing_variables = self._create_all_conversation_variables(session)
else:
# Check and add any missing variables from the workflow
existing_variables = self._sync_missing_conversation_variables(session, existing_variables)
existing_variables = self._sync_missing_conversation_variables(existing_variables, session)
# Convert to Variable objects for use in the workflow
conversation_variables = [var.to_variable() for var in existing_variables]
@@ -435,7 +446,7 @@ class AdvancedChatAppRunner(WorkflowBasedAppRunner):
return new_variables
def _sync_missing_conversation_variables(
self, session: Session, existing_variables: list[ConversationVariable]
self, existing_variables: list[ConversationVariable], session: Session
) -> list[ConversationVariable]:
"""
Sync missing conversation variables from the workflow definition.
@@ -9,7 +9,7 @@ from threading import Thread
from typing import Any, Union
from sqlalchemy import select, update
from sqlalchemy.orm import Session, sessionmaker
from sqlalchemy.orm import Session
from constants.tts_auto_play_timeout import TTS_AUTO_PLAY_TIMEOUT, TTS_AUTO_PLAY_YIELD_CPU_TIME
from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom
@@ -71,12 +71,12 @@ from core.app.entities.task_entities import (
from core.app.task_pipeline.based_generate_task_pipeline import BasedGenerateTaskPipeline
from core.app.task_pipeline.message_cycle_manager import MessageCycleManager
from core.base.tts import AppGeneratorTTSPublisher, AudioTrunk
from core.db.session_factory import session_factory
from core.ops.ops_trace_manager import TraceQueueManager
from core.repositories.human_input_repository import HumanInputFormRepositoryImpl
from core.workflow.file_reference import resolve_file_record_id
from core.workflow.nodes.human_input.pause_reason import HumanInputRequired
from core.workflow.system_variables import build_system_variables
from extensions.ext_database import db
from graphon.enums import WorkflowExecutionStatus
from graphon.model_runtime.entities.llm_entities import LLMUsage
from graphon.model_runtime.utils.encoders import jsonable_encoder
@@ -399,8 +399,13 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
@contextmanager
def _database_session(self):
"""Context manager for database sessions."""
with sessionmaker(bind=db.engine, expire_on_commit=False).begin() as session:
yield session
with session_factory.create_session() as session:
try:
yield session
session.commit()
except Exception:
session.rollback()
raise
def _ensure_workflow_initialized(self):
"""Fluent validation for workflow state."""
@@ -825,7 +830,8 @@ class AdvancedChatAppGenerateTaskPipeline(GraphRuntimeStateSupport):
self, event: QueueAnnotationReplyEvent, **kwargs
) -> Generator[StreamResponse, None, None]:
"""Handle annotation reply events."""
self._message_cycle_manager.handle_annotation_reply(event)
with self._database_session() as session:
self._message_cycle_manager.handle_annotation_reply(event, session)
yield from ()
def _handle_message_replace_event(
@@ -24,7 +24,7 @@ from core.app.app_config.entities import (
from core.app.apps.agent_app.app_feature_projection import merge_agent_app_features
from core.app.apps.agent_app.app_variable_projection import agent_app_variables_to_user_input_form
from models.agent_config_entities import AgentSoulConfig
from models.model import App, AppMode, AppModelConfig, AppModelConfigDict, Conversation
from models.model import AnnotationReplyConfig, App, AppMode, AppModelConfig, AppModelConfigDict, Conversation
class AgentAppConfig(EasyUIBasedAppConfig):
@@ -43,11 +43,16 @@ class AgentAppConfigManager(BaseAppConfigManager):
*,
app_model: App,
agent_soul: AgentSoulConfig,
annotation_reply: AnnotationReplyConfig | None,
app_model_config: AppModelConfig | None = None,
conversation: Conversation | None = None,
) -> AgentAppConfig:
"""Build the Agent App config from the Agent Soul (+ optional feature flags)."""
config_dict = cls._synthesize_config_dict(agent_soul, app_model_config)
config_dict = cls._synthesize_config_dict(
agent_soul,
app_model_config,
annotation_reply=annotation_reply,
)
# The synthesized dict is shaped like an app_model_config; the EasyUI
# sub-managers type their param as AppModelConfigDict (a TypedDict).
typed_config = cast(AppModelConfigDict, config_dict)
@@ -77,6 +82,8 @@ class AgentAppConfigManager(BaseAppConfigManager):
def _synthesize_config_dict(
agent_soul: AgentSoulConfig,
app_model_config: AppModelConfig | None,
*,
annotation_reply: AnnotationReplyConfig | None,
) -> dict[str, Any]:
"""Shape a Soul + feature flags into an ``app_model_config``-style dict.
@@ -84,7 +91,11 @@ class AgentAppConfigManager(BaseAppConfigManager):
``app_model_config`` when one exists; model + prompt always come from
the Agent Soul (the single source of truth for those).
"""
base = merge_agent_app_features(agent_soul=agent_soul, app_model_config=app_model_config)
base = merge_agent_app_features(
agent_soul=agent_soul,
app_model_config=app_model_config,
annotation_reply=annotation_reply,
)
model = agent_soul.model
if model is not None:
@@ -1,12 +1,14 @@
from typing import Any
from models.agent_config_entities import AgentSoulConfig
from models.model import AnnotationReplyConfig
def merge_agent_app_features(
*,
agent_soul: AgentSoulConfig,
app_model_config: Any | None,
annotation_reply: AnnotationReplyConfig | None,
) -> dict[str, Any]:
"""Project public Agent App features from legacy config plus Agent Soul.
@@ -14,7 +16,12 @@ def merge_agent_app_features(
opening statements. Agent Soul is the source of truth for Agent-owned
features like file upload, so Soul fields override same-named legacy keys.
"""
features: dict[str, Any] = dict(app_model_config.to_dict()) if app_model_config else {}
if app_model_config is None:
features: dict[str, Any] = {}
else:
if annotation_reply is None:
raise ValueError("Annotation reply config is required")
features = dict(app_model_config.to_dict(annotation_reply=annotation_reply))
soul_features = agent_soul.app_features.model_dump(mode="json", exclude_none=True)
features.update(soul_features)
return features
+93 -53
View File
@@ -21,6 +21,7 @@ from typing import Any, Literal
from flask import Flask, current_app
from pydantic import JsonValue
from sqlalchemy import and_, or_, select
from sqlalchemy.orm import Session
from clients.agent_backend import AgentBackendRunEventAdapter
from clients.agent_backend.factory import create_agent_backend_run_client
@@ -46,10 +47,11 @@ from core.app.entities.app_invoke_entities import (
UserFrom,
)
from core.app.llm.model_access import build_dify_model_access
from core.db.session_factory import session_factory
from core.ops.ops_trace_manager import TraceQueueManager
from core.workflow.file_reference import build_file_reference, is_canonical_file_reference
from extensions.ext_database import db
from models import Account, App, EndUser, Message
from models import Account, App, AppModelConfig, EndUser, Message, MessageAnnotation
from models.agent import (
APP_BACKED_AGENT_SOURCES,
Agent,
@@ -61,6 +63,7 @@ from models.agent import (
AgentStatus,
)
from models.agent_config_entities import AgentSoulConfig
from models.model import load_annotation_reply_config
from services.conversation_service import ConversationService
logger = logging.getLogger(__name__)
@@ -137,6 +140,7 @@ class AgentAppGenerator(MessageBasedAppGenerator):
user: Account | EndUser,
args: Mapping[str, Any],
invoke_from: InvokeFrom,
session: Session,
streaming: bool = True,
) -> Mapping[str, Any] | Generator[Mapping | str, None, None]:
if not streaming:
@@ -152,6 +156,7 @@ class AgentAppGenerator(MessageBasedAppGenerator):
invoke_from=invoke_from,
draft_type=args.get("draft_type"),
user=user,
session=session,
)
runtime_session_snapshot_id = self._runtime_session_snapshot_id(
invoke_from=invoke_from,
@@ -162,15 +167,19 @@ class AgentAppGenerator(MessageBasedAppGenerator):
conversation_id = args.get("conversation_id")
if conversation_id:
conversation = ConversationService.get_conversation(
app_model=app_model, conversation_id=conversation_id, user=user, session=db.session()
app_model=app_model, conversation_id=conversation_id, user=user, session=session
)
# Build the EasyUI-shaped config from the Agent Soul so the chat pipeline
# can persist usage; the answer itself comes from the agent backend.
app_model_config = app_model.app_model_config
app_model_config = (
session.get(AppModelConfig, app_model.app_model_config_id) if app_model.app_model_config_id else None
)
annotation_reply = load_annotation_reply_config(session, app_model.id) if app_model_config else None
app_config = AgentAppConfigManager.get_app_config(
app_model=app_model,
agent_soul=agent_soul,
annotation_reply=annotation_reply,
app_model_config=app_model_config,
conversation=conversation,
)
@@ -210,7 +219,11 @@ class AgentAppGenerator(MessageBasedAppGenerator):
agent_runtime_exit_intent=agent_runtime_exit_intent,
)
conversation, message = self._init_generate_records(application_generate_entity, conversation)
conversation, message = self._init_generate_records(
application_generate_entity,
conversation,
session=session,
)
queue_manager = MessageBasedAppQueueManager(
task_id=application_generate_entity.task_id,
@@ -253,6 +266,7 @@ class AgentAppGenerator(MessageBasedAppGenerator):
user: Account | EndUser,
conversation_id: str,
invoke_from: InvokeFrom,
session: Session,
) -> None:
"""Resume an Agent App conversation after a submitted ask_human HITL form.
@@ -263,19 +277,28 @@ class AgentAppGenerator(MessageBasedAppGenerator):
out of scope here — the message is persisted and can be re-fetched.
"""
conversation = ConversationService.get_conversation(
app_model=app_model, conversation_id=conversation_id, user=user, session=db.session()
app_model=app_model, conversation_id=conversation_id, user=user, session=session
)
agent, agent_config_id, agent_config_version_kind, agent_soul = self._resolve_agent(
app_model,
invoke_from=invoke_from,
draft_type=self._resume_draft_type(app_model=app_model, conversation=conversation, user=user),
draft_type=self._resume_draft_type(
app_model=app_model, conversation=conversation, user=user, session=session
),
user=user,
session=session,
)
app_model_config = (
session.get(AppModelConfig, app_model.app_model_config_id) if app_model.app_model_config_id else None
)
annotation_reply = load_annotation_reply_config(session, app_model.id) if app_model_config else None
app_config = AgentAppConfigManager.get_app_config(
app_model=app_model,
agent_soul=agent_soul,
app_model_config=app_model.app_model_config,
annotation_reply=annotation_reply,
app_model_config=app_model_config,
conversation=conversation,
)
model_conf = ModelConfigConverter.convert(app_config)
@@ -287,7 +310,7 @@ class AgentAppGenerator(MessageBasedAppGenerator):
# turn's query); the continuation is driven by deferred_tool_results and
# the restored snapshot, not by re-processing this prompt. A blank prompt
# would drop the user-prompt layer and fail the snapshot match.
paused_message = db.session.scalar(
paused_message = session.scalar(
select(Message)
.where(Message.conversation_id == conversation.id, Message.query != "")
.order_by(Message.created_at.desc())
@@ -318,7 +341,11 @@ class AgentAppGenerator(MessageBasedAppGenerator):
agent_config_version_kind=agent_config_version_kind,
)
conversation, message = self._init_generate_records(application_generate_entity, conversation)
conversation, message = self._init_generate_records(
application_generate_entity,
conversation,
session=session,
)
queue_manager = MessageBasedAppQueueManager(
task_id=application_generate_entity.task_id,
@@ -357,7 +384,9 @@ class AgentAppGenerator(MessageBasedAppGenerator):
)
@staticmethod
def _resume_draft_type(*, app_model: App, conversation: Any, user: Account | EndUser) -> str | None:
def _resume_draft_type(
*, app_model: App, conversation: Any, user: Account | EndUser, session: Session
) -> str | None:
if conversation.invoke_from != InvokeFrom.DEBUGGER:
return None
active_session = AgentAppRuntimeSessionStore().load_active_session_for_conversation(
@@ -367,7 +396,7 @@ class AgentAppGenerator(MessageBasedAppGenerator):
)
snapshot_id = active_session.scope.agent_config_snapshot_id if active_session is not None else None
if snapshot_id and isinstance(user, Account):
draft = db.session.scalar(
draft = session.scalar(
select(AgentConfigDraft).where(
AgentConfigDraft.tenant_id == app_model.tenant_id,
AgentConfigDraft.id == snapshot_id,
@@ -413,15 +442,31 @@ class AgentAppGenerator(MessageBasedAppGenerator):
# Apply app-level input guards (content moderation + annotation
# reply) before reaching the Agent backend, mirroring the EasyUI
# chat / agent-chat runners. These can short-circuit the turn.
app_model = db.session.get(App, app_config.app_id)
if app_model is None:
raise AgentAppGeneratorError("App not found")
handled, query = self._run_input_guards(
application_generate_entity=application_generate_entity,
app_model=app_model,
message=message,
queue_manager=queue_manager,
)
with session_factory.get_session_maker().begin() as session:
app_model = session.get(App, app_config.app_id)
if app_model is None:
raise AgentAppGeneratorError("App not found")
handled, query, annotation_reply = self._run_input_guards(
session=session,
application_generate_entity=application_generate_entity,
app_model=app_model,
message=message,
queue_manager=queue_manager,
)
if annotation_reply:
from core.app.apps.agent_app.app_runner import publish_text_answer
from core.app.entities.queue_entities import QueueAnnotationReplyEvent
queue_manager.publish(
QueueAnnotationReplyEvent(message_annotation_id=annotation_reply.id),
PublishFrom.APPLICATION_MANAGER,
)
publish_text_answer(
queue_manager=queue_manager,
model_name=application_generate_entity.model_conf.model,
answer=annotation_reply.content,
user_query=query,
)
if handled:
return
query = _append_prompt_file_mappings(
@@ -436,11 +481,13 @@ class AgentAppGenerator(MessageBasedAppGenerator):
user_from=user_from,
invoke_from=application_generate_entity.invoke_from,
)
_, _, agent_soul = self._resolve_agent_by_id(
tenant_id=app_config.tenant_id,
agent_id=application_generate_entity.agent_id,
snapshot_id=application_generate_entity.agent_config_snapshot_id,
)
with session_factory.create_session() as session:
_, _, agent_soul = self._resolve_agent_by_id(
tenant_id=app_config.tenant_id,
agent_id=application_generate_entity.agent_id,
snapshot_id=application_generate_entity.agent_config_snapshot_id,
session=session,
)
runner = self._build_runner(dify_context)
runner.run(
@@ -502,20 +549,18 @@ class AgentAppGenerator(MessageBasedAppGenerator):
def _run_input_guards(
self,
*,
session: Session,
application_generate_entity: AgentAppGenerateEntity,
app_model: App,
message: Message,
queue_manager: AppQueueManager,
) -> tuple[bool, str]:
) -> tuple[bool, str, MessageAnnotation | None]:
"""Apply input moderation + annotation reply before the backend call.
Returns ``(handled, query)``: when ``handled`` is True a direct answer
has already been published (a blocked/preset moderation response or a
matched annotation) and the backend turn must be skipped. Otherwise
``query`` is the possibly moderation-overridden query to send onward.
Returns ``(handled, query, annotation_reply)``. Annotation output is
published by the caller only after this transaction commits.
"""
from core.app.apps.agent_app.app_runner import publish_text_answer
from core.app.entities.queue_entities import QueueAnnotationReplyEvent
from core.app.features.annotation_reply.annotation_reply import AnnotationReplyFeature
from core.moderation.base import ModerationError
from core.moderation.input_moderation import InputModeration
@@ -538,7 +583,7 @@ class AgentAppGenerator(MessageBasedAppGenerator):
)
except ModerationError as e:
publish_text_answer(queue_manager=queue_manager, model_name=model_name, answer=str(e), user_query=query)
return True, query
return True, query, None
# annotation reply: a matching annotation answers the turn deterministically.
if query:
@@ -548,21 +593,12 @@ class AgentAppGenerator(MessageBasedAppGenerator):
query=query,
user_id=application_generate_entity.user_id,
invoke_from=application_generate_entity.invoke_from,
session=session,
)
if annotation_reply:
queue_manager.publish(
QueueAnnotationReplyEvent(message_annotation_id=annotation_reply.id),
PublishFrom.APPLICATION_MANAGER,
)
publish_text_answer(
queue_manager=queue_manager,
model_name=model_name,
answer=annotation_reply.content,
user_query=query,
)
return True, query
return True, query, annotation_reply
return False, query
return False, query, None
def _resolve_agent(
self,
@@ -571,8 +607,9 @@ class AgentAppGenerator(MessageBasedAppGenerator):
invoke_from: InvokeFrom,
draft_type: Any,
user: Account | EndUser,
session: Session,
) -> tuple[Agent, str, Literal["snapshot", "draft", "build_draft"], AgentSoulConfig]:
agent = db.session.scalar(
agent = session.scalar(
select(Agent)
.where(
Agent.tenant_id == app_model.tenant_id,
@@ -603,6 +640,7 @@ class AgentAppGenerator(MessageBasedAppGenerator):
agent=agent,
draft_type=draft_type,
account_id=user.id if isinstance(user, Account) else None,
session=session,
)
agent_soul = AgentSoulConfig.model_validate(draft.config_snapshot_dict)
config_version_kind: Literal["snapshot", "draft", "build_draft"] = (
@@ -617,6 +655,7 @@ class AgentAppGenerator(MessageBasedAppGenerator):
tenant_id=app_model.tenant_id,
agent_id=agent.id,
snapshot_id=agent.active_config_snapshot_id,
session=session,
)
return agent, snapshot.id, "snapshot", agent_soul
@@ -633,7 +672,7 @@ class AgentAppGenerator(MessageBasedAppGenerator):
@staticmethod
def _resolve_debug_draft(
*, tenant_id: str, agent: Agent, draft_type: Any, account_id: str | None
*, tenant_id: str, agent: Agent, draft_type: Any, account_id: str | None, session: Session
) -> AgentConfigDraft:
effective_draft_type = (
AgentConfigDraftType.DEBUG_BUILD
@@ -651,7 +690,7 @@ class AgentAppGenerator(MessageBasedAppGenerator):
stmt = stmt.where(AgentConfigDraft.account_id == account_id)
else:
stmt = stmt.where(AgentConfigDraft.account_id.is_(None))
draft = db.session.scalar(stmt.order_by(AgentConfigDraft.updated_at.desc()).limit(1))
draft = session.scalar(stmt.order_by(AgentConfigDraft.updated_at.desc()).limit(1))
if draft is not None:
return draft
if effective_draft_type == AgentConfigDraftType.DEBUG_BUILD:
@@ -660,6 +699,7 @@ class AgentAppGenerator(MessageBasedAppGenerator):
tenant_id=tenant_id,
agent_id=agent.id,
snapshot_id=agent.active_config_snapshot_id,
session=session,
)
draft = AgentConfigDraft(
tenant_id=tenant_id,
@@ -672,20 +712,20 @@ class AgentAppGenerator(MessageBasedAppGenerator):
created_by=agent.created_by,
updated_by=agent.updated_by,
)
db.session.add(draft)
db.session.flush()
session.add(draft)
session.flush()
return draft
@staticmethod
def _resolve_agent_by_id(
*, tenant_id: str, agent_id: str, snapshot_id: str | None
*, tenant_id: str, agent_id: str, snapshot_id: str | None, session: Session
) -> tuple[Agent, AgentConfigSnapshot | AgentConfigDraft, AgentSoulConfig]:
agent = db.session.scalar(select(Agent).where(Agent.id == agent_id, Agent.tenant_id == tenant_id))
agent = session.scalar(select(Agent).where(Agent.id == agent_id, Agent.tenant_id == tenant_id))
if agent is None:
raise AgentAppGeneratorError("Agent not found")
if not snapshot_id:
raise AgentAppGeneratorError("Agent has no published version")
snapshot = db.session.scalar(
snapshot = session.scalar(
select(AgentConfigSnapshot).where(
AgentConfigSnapshot.tenant_id == tenant_id,
AgentConfigSnapshot.agent_id == agent_id,
@@ -695,7 +735,7 @@ class AgentAppGenerator(MessageBasedAppGenerator):
if snapshot is not None:
agent_soul = AgentSoulConfig.model_validate(snapshot.config_snapshot_dict)
return agent, snapshot, agent_soul
draft = db.session.scalar(
draft = session.scalar(
select(AgentConfigDraft).where(
AgentConfigDraft.tenant_id == tenant_id,
AgentConfigDraft.agent_id == agent_id,
@@ -2,6 +2,8 @@ import uuid
from collections.abc import Mapping
from typing import Any, cast
from sqlalchemy.orm import Session
from core.agent.entities import AgentEntity
from core.app.app_config.base_app_config_manager import BaseAppConfigManager
from core.app.app_config.common.sensitive_word_avoidance.manager import SensitiveWordAvoidanceConfigManager
@@ -20,7 +22,7 @@ from core.app.app_config.features.suggested_questions_after_answer.manager impor
)
from core.app.app_config.features.text_to_speech.manager import TextToSpeechConfigManager
from core.entities.agent_entities import PlanningStrategy
from models.model import App, AppMode, AppModelConfig, AppModelConfigDict, Conversation
from models.model import AnnotationReplyConfig, App, AppMode, AppModelConfig, AppModelConfigDict, Conversation
OLD_TOOLS = ["dataset", "google_search", "web_reader", "wikipedia", "current_datetime"]
@@ -41,6 +43,8 @@ class AgentChatAppConfigManager(BaseAppConfigManager):
app_model_config: AppModelConfig,
conversation: Conversation | None = None,
override_config_dict: AppModelConfigDict | None = None,
*,
annotation_reply: AnnotationReplyConfig | None,
) -> AgentChatAppConfig:
"""
Convert app model config to agent chat app config
@@ -58,7 +62,7 @@ class AgentChatAppConfigManager(BaseAppConfigManager):
config_from = EasyUIBasedAppModelConfigFrom.APP_LATEST_CONFIG
if config_from != EasyUIBasedAppModelConfigFrom.ARGS:
app_model_config_dict = app_model_config.to_dict()
app_model_config_dict = app_model_config.to_dict(annotation_reply=annotation_reply)
config_dict = app_model_config_dict.copy()
else:
if not override_config_dict:
@@ -88,7 +92,7 @@ class AgentChatAppConfigManager(BaseAppConfigManager):
return app_config
@classmethod
def config_validate(cls, tenant_id: str, config: Mapping[str, Any]) -> AppModelConfigDict:
def config_validate(cls, tenant_id: str, config: Mapping[str, Any], session: Session) -> AppModelConfigDict:
"""
Validate for agent chat app model config
@@ -116,7 +120,7 @@ class AgentChatAppConfigManager(BaseAppConfigManager):
related_config_keys.extend(current_related_config_keys)
# agent_mode
config, current_related_config_keys = cls.validate_agent_mode_and_set_defaults(tenant_id, config)
config, current_related_config_keys = cls.validate_agent_mode_and_set_defaults(tenant_id, config, session)
related_config_keys.extend(current_related_config_keys)
# opening_statement
@@ -144,7 +148,7 @@ class AgentChatAppConfigManager(BaseAppConfigManager):
# dataset configs
# dataset_query_variable
config, current_related_config_keys = DatasetConfigManager.validate_and_set_defaults(
tenant_id, app_mode, config
tenant_id, app_mode, config, session
)
related_config_keys.extend(current_related_config_keys)
@@ -163,7 +167,7 @@ class AgentChatAppConfigManager(BaseAppConfigManager):
@classmethod
def validate_agent_mode_and_set_defaults(
cls, tenant_id: str, config: dict[str, Any]
cls, tenant_id: str, config: dict[str, Any], session: Session
) -> tuple[dict[str, Any], list[str]]:
"""
Validate agent_mode and set defaults for agent feature
@@ -220,7 +224,7 @@ class AgentChatAppConfigManager(BaseAppConfigManager):
except ValueError:
raise ValueError("id in dataset must be of UUID type")
if not DatasetConfigManager.is_dataset_exists(tenant_id, tool_item["id"]):
if not DatasetConfigManager.is_dataset_exists(tenant_id, tool_item["id"], session):
raise ValueError("Dataset ID does not exist, please check your permission.")
else:
# latest style, use key-value pair
+32 -16
View File
@@ -21,13 +21,14 @@ from core.app.apps.exc import GenerateTaskStoppedError
from core.app.apps.message_based_app_generator import MessageBasedAppGenerator
from core.app.apps.message_based_app_queue_manager import MessageBasedAppQueueManager
from core.app.entities.app_invoke_entities import AgentChatAppGenerateEntity, InvokeFrom
from core.db.session_factory import session_factory
from core.helper.trace_id_helper import extract_trace_session_id_from_args
from core.ops.ops_trace_manager import TraceQueueManager
from extensions.ext_database import db
from factories import file_factory
from graphon.model_runtime.errors.invoke import InvokeAuthorizationError
from libs.flask_utils import preserve_flask_contexts
from models import Account, App, EndUser
from models.model import load_annotation_reply_config
from services.conversation_service import ConversationService
logger = logging.getLogger(__name__)
@@ -43,6 +44,7 @@ class AgentChatAppGenerator(MessageBasedAppGenerator):
args: Mapping[str, Any],
invoke_from: InvokeFrom,
streaming: Literal[False],
session: Session,
) -> Mapping[str, Any]: ...
@overload
@@ -54,6 +56,7 @@ class AgentChatAppGenerator(MessageBasedAppGenerator):
args: Mapping[str, Any],
invoke_from: InvokeFrom,
streaming: Literal[True],
session: Session,
) -> Generator[Mapping | str, None, None]: ...
@overload
@@ -65,6 +68,7 @@ class AgentChatAppGenerator(MessageBasedAppGenerator):
args: Mapping[str, Any],
invoke_from: InvokeFrom,
streaming: bool,
session: Session,
) -> Mapping | Generator[Mapping | str, None, None]: ...
def generate(
@@ -75,6 +79,7 @@ class AgentChatAppGenerator(MessageBasedAppGenerator):
args: Mapping[str, Any],
invoke_from: InvokeFrom,
streaming: bool = True,
session: Session,
) -> Mapping | Generator[Mapping | str, None, None]:
"""
Generate App response.
@@ -108,10 +113,14 @@ class AgentChatAppGenerator(MessageBasedAppGenerator):
conversation_id = args.get("conversation_id")
if conversation_id:
conversation = ConversationService.get_conversation(
app_model=app_model, conversation_id=conversation_id, user=user, session=db.session()
app_model=app_model, conversation_id=conversation_id, user=user, session=session
)
# get app model config
app_model_config = self._get_app_model_config(app_model=app_model, conversation=conversation)
app_model_config = self._get_app_model_config(
app_model=app_model,
conversation=conversation,
session=session,
)
# validate override model config
override_model_config_dict = None
@@ -123,11 +132,16 @@ class AgentChatAppGenerator(MessageBasedAppGenerator):
override_model_config_dict = AgentChatAppConfigManager.config_validate(
tenant_id=app_model.tenant_id,
config=args["model_config"],
session=session,
)
# always enable retriever resource in debugger mode
override_model_config_dict["retriever_resource"] = {"enabled": True}
annotation_reply = (
None if override_model_config_dict else load_annotation_reply_config(session, app_model_config.app_id)
)
# parse files
# TODO(QuantumGhost): Move file parsing logic to the API controller layer
# for better separation of concerns.
@@ -137,7 +151,7 @@ class AgentChatAppGenerator(MessageBasedAppGenerator):
with self._bind_file_access_scope(tenant_id=app_model.tenant_id, user=user, invoke_from=invoke_from):
files = args.get("files") or []
file_extra_config = FileUploadConfigManager.convert(
override_model_config_dict or app_model_config.to_dict()
override_model_config_dict or app_model_config.to_dict(annotation_reply=annotation_reply)
)
if file_extra_config:
file_objs = file_factory.build_from_mappings(
@@ -155,6 +169,7 @@ class AgentChatAppGenerator(MessageBasedAppGenerator):
app_model_config=app_model_config,
conversation=conversation,
override_config_dict=override_model_config_dict,
annotation_reply=annotation_reply,
)
# get tracing instance
@@ -186,7 +201,11 @@ class AgentChatAppGenerator(MessageBasedAppGenerator):
)
# init generate records
(conversation, message) = self._init_generate_records(application_generate_entity, conversation)
(conversation, message) = self._init_generate_records(
application_generate_entity,
conversation,
session=session,
)
# init queue manager
queue_manager = MessageBasedAppQueueManager(
@@ -205,7 +224,6 @@ class AgentChatAppGenerator(MessageBasedAppGenerator):
target=self._generate_worker,
kwargs={
"flask_app": current_app._get_current_object(), # type: ignore
"session": db.session(),
"context": context,
"application_generate_entity": application_generate_entity,
"queue_manager": queue_manager,
@@ -230,7 +248,6 @@ class AgentChatAppGenerator(MessageBasedAppGenerator):
def _generate_worker(
self,
flask_app: Flask,
session: Session,
context: contextvars.Context,
application_generate_entity: AgentChatAppGenerateEntity,
queue_manager: AppQueueManager,
@@ -255,13 +272,14 @@ class AgentChatAppGenerator(MessageBasedAppGenerator):
# chatbot app
runner = AgentChatAppRunner()
runner.run(
session=session,
application_generate_entity=application_generate_entity,
queue_manager=queue_manager,
conversation=conversation,
message=message,
)
with session_factory.create_session() as session:
runner.run(
application_generate_entity=application_generate_entity,
queue_manager=queue_manager,
conversation=conversation,
message=message,
session=session,
)
except GenerateTaskStoppedError:
pass
except InvokeAuthorizationError:
@@ -278,5 +296,3 @@ class AgentChatAppGenerator(MessageBasedAppGenerator):
except Exception as e:
logger.exception("Unknown Error when generating")
queue_manager.publish_error(e, PublishFrom.APPLICATION_MANAGER)
finally:
db.session.close()
+20 -11
View File
@@ -32,14 +32,17 @@ class AgentChatAppRunner(AppRunner):
def run(
self,
session: Session,
application_generate_entity: AgentChatAppGenerateEntity,
queue_manager: AppQueueManager,
conversation: Conversation,
message: Message,
session: Session,
):
"""
Run assistant application
"""Run the assistant application with bounded explicit transactions.
The setup session is committed and released before the multi-step agent
runner begins model or tool I/O.
:param application_generate_entity: application generate entity
:param queue_manager: application queue manager
:param conversation: conversation
@@ -49,10 +52,10 @@ class AgentChatAppRunner(AppRunner):
app_config = application_generate_entity.app_config
app_config = cast(AgentChatAppConfig, app_config)
app_stmt = select(App).where(App.id == app_config.app_id)
with create_session() as session:
app_record = session.scalar(app_stmt)
with create_session() as read_session:
app_record = read_session.scalar(app_stmt)
if app_record:
session.expunge(app_record)
read_session.expunge(app_record)
if not app_record:
raise ValueError("App not found")
@@ -112,7 +115,10 @@ class AgentChatAppRunner(AppRunner):
query=query,
user_id=application_generate_entity.user_id,
invoke_from=application_generate_entity.invoke_from,
session=session,
)
session.commit()
session.close()
if annotation_reply:
queue_manager.publish(
@@ -191,15 +197,15 @@ class AgentChatAppRunner(AppRunner):
agent_entity.strategy = AgentEntity.Strategy.FUNCTION_CALLING
conversation_stmt = select(Conversation).where(Conversation.id == conversation.id)
msg_stmt = select(Message).where(Message.id == message.id)
with create_session() as session:
conversation_result = session.scalar(conversation_stmt)
with create_session() as read_session:
conversation_result = read_session.scalar(conversation_stmt)
if conversation_result is None:
raise ValueError("Conversation not found")
message_result = session.scalar(msg_stmt)
message_result = read_session.scalar(msg_stmt)
if message_result is not None:
session.expunge(message_result)
session.expunge(conversation_result)
read_session.expunge(message_result)
read_session.expunge(conversation_result)
if message_result is None:
raise ValueError("Message not found")
@@ -234,6 +240,9 @@ class AgentChatAppRunner(AppRunner):
model_instance=model_instance,
)
session.commit()
session.close()
invoke_result = runner.run(
session=session,
message=message,
+36 -27
View File
@@ -5,7 +5,7 @@ from collections.abc import Generator, Mapping, Sequence
from mimetypes import guess_extension
from typing import TYPE_CHECKING, Any, Union
from sqlalchemy.orm import sessionmaker
from sqlalchemy.orm import Session
from core.app.app_config.entities import ExternalDataVariableEntity, PromptTemplateEntity
from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom
@@ -24,6 +24,7 @@ from core.app.entities.queue_entities import (
)
from core.app.features.annotation_reply.annotation_reply import AnnotationReplyFeature
from core.app.features.hosting_moderation.hosting_moderation import HostingModerationFeature
from core.db.session_factory import session_factory
from core.external_data_tool.external_data_fetch import ExternalDataFetch
from core.memory.token_buffer_memory import TokenBufferMemory
from core.model_manager import ModelInstance
@@ -32,7 +33,6 @@ from core.prompt.advanced_prompt_transform import AdvancedPromptTransform
from core.prompt.entities.advanced_prompt_entities import ChatModelMessage, CompletionModelPromptTemplate, MemoryConfig
from core.prompt.simple_prompt_transform import ModelMode, SimplePromptTransform
from core.tools.tool_file_manager import ToolFileManager
from extensions.ext_database import db
from graphon.file import FileTransferMethod, FileType
from graphon.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk, LLMResultChunkDelta, LLMUsage
from graphon.model_runtime.entities.message_entities import (
@@ -310,20 +310,30 @@ class AppRunner:
case list():
for content in message.content:
match content:
case str():
text += content
case TextPromptMessageContent():
text += content.data
case ImagePromptMessageContent():
if message_id and user_id and tenant_id:
try:
self._handle_multimodal_image_content(
content=content,
message_id=message_id,
user_id=user_id,
tenant_id=tenant_id,
queue_manager=queue_manager,
)
with session_factory.create_session() as session:
message_file_id = self._handle_multimodal_image_content(
session=session,
content=content,
message_id=message_id,
user_id=user_id,
tenant_id=tenant_id,
queue_manager=queue_manager,
)
session.commit()
if message_file_id:
queue_manager.publish(
QueueMessageFileEvent(message_file_id=message_file_id),
PublishFrom.APPLICATION_MANAGER,
)
_logger.info(
"QueueMessageFileEvent published for message_file_id: %s",
message_file_id,
)
except Exception:
_logger.exception("Failed to handle multimodal image output")
else:
@@ -365,7 +375,8 @@ class AppRunner:
user_id: str,
tenant_id: str,
queue_manager: AppQueueManager,
):
session: Session,
) -> str | None:
"""
Handle multimodal image content from LLM response.
Save the image and create a MessageFile record.
@@ -386,7 +397,7 @@ class AppRunner:
if not image_url and not base64_data:
_logger.warning("Image content has neither URL nor base64 data")
return
return None
tool_file_manager = ToolFileManager()
@@ -420,10 +431,10 @@ class AppRunner:
)
_logger.info("Image saved successfully, tool_file_id: %s", tool_file.id)
else:
return
return None
except Exception:
_logger.exception("Failed to save image file")
return
return None
# Create MessageFile record.
# Use an independent session so this side-effect write does not
@@ -441,16 +452,9 @@ class AppRunner:
created_by=user_id,
)
with sessionmaker(bind=db.engine, expire_on_commit=False).begin() as session:
session.add(message_file)
# Publish QueueMessageFileEvent
queue_manager.publish(
QueueMessageFileEvent(message_file_id=message_file.id),
PublishFrom.APPLICATION_MANAGER,
)
_logger.info("QueueMessageFileEvent published for message_file_id: %s", message_file.id)
session.add(message_file)
session.flush()
return message_file.id
def moderation_for_inputs(
self,
@@ -536,7 +540,7 @@ class AppRunner:
)
def query_app_annotations_to_reply(
self, app_record: App, message: Message, query: str, user_id: str, invoke_from: InvokeFrom
self, app_record: App, message: Message, query: str, user_id: str, invoke_from: InvokeFrom, session: Session
) -> MessageAnnotation | None:
"""
Query app annotations to reply
@@ -549,5 +553,10 @@ class AppRunner:
"""
annotation_reply_feature = AnnotationReplyFeature()
return annotation_reply_feature.query(
app_record=app_record, message=message, query=query, user_id=user_id, invoke_from=invoke_from
app_record=app_record,
message=message,
query=query,
user_id=user_id,
invoke_from=invoke_from,
session=session,
)
+8 -4
View File
@@ -1,5 +1,7 @@
from typing import Any, cast
from sqlalchemy.orm import Session
from core.app.app_config.base_app_config_manager import BaseAppConfigManager
from core.app.app_config.common.sensitive_word_avoidance.manager import SensitiveWordAvoidanceConfigManager
from core.app.app_config.easy_ui_based_app.dataset.manager import DatasetConfigManager
@@ -15,7 +17,7 @@ from core.app.app_config.features.suggested_questions_after_answer.manager impor
SuggestedQuestionsAfterAnswerConfigManager,
)
from core.app.app_config.features.text_to_speech.manager import TextToSpeechConfigManager
from models.model import App, AppMode, AppModelConfig, AppModelConfigDict, Conversation
from models.model import AnnotationReplyConfig, App, AppMode, AppModelConfig, AppModelConfigDict, Conversation
class ChatAppConfig(EasyUIBasedAppConfig):
@@ -34,6 +36,8 @@ class ChatAppConfigManager(BaseAppConfigManager):
app_model_config: AppModelConfig,
conversation: Conversation | None = None,
override_config_dict: AppModelConfigDict | None = None,
*,
annotation_reply: AnnotationReplyConfig | None,
) -> ChatAppConfig:
"""
Convert app model config to chat app config
@@ -51,7 +55,7 @@ class ChatAppConfigManager(BaseAppConfigManager):
config_from = EasyUIBasedAppModelConfigFrom.APP_LATEST_CONFIG
if config_from != EasyUIBasedAppModelConfigFrom.ARGS:
app_model_config_dict = app_model_config.to_dict()
app_model_config_dict = app_model_config.to_dict(annotation_reply=annotation_reply)
config_dict = app_model_config_dict.copy()
else:
if not override_config_dict:
@@ -81,7 +85,7 @@ class ChatAppConfigManager(BaseAppConfigManager):
return app_config
@classmethod
def config_validate(cls, tenant_id: str, config: dict[str, Any]) -> AppModelConfigDict:
def config_validate(cls, tenant_id: str, config: dict[str, Any], session: Session) -> AppModelConfigDict:
"""
Validate for chat app model config
@@ -110,7 +114,7 @@ class ChatAppConfigManager(BaseAppConfigManager):
# dataset_query_variable
config, current_related_config_keys = DatasetConfigManager.validate_and_set_defaults(
tenant_id, app_mode, config
tenant_id, app_mode, config, session=session
)
related_config_keys.extend(current_related_config_keys)
+36 -19
View File
@@ -21,13 +21,14 @@ from core.app.apps.exc import GenerateTaskStoppedError
from core.app.apps.message_based_app_generator import MessageBasedAppGenerator
from core.app.apps.message_based_app_queue_manager import MessageBasedAppQueueManager
from core.app.entities.app_invoke_entities import ChatAppGenerateEntity, InvokeFrom
from core.db.session_factory import session_factory
from core.helper.trace_id_helper import extract_trace_session_id_from_args
from core.ops.ops_trace_manager import TraceQueueManager
from extensions.ext_database import db
from factories import file_factory
from graphon.model_runtime.errors.invoke import InvokeAuthorizationError
from models import Account
from models.model import App, EndUser
from models.model import App, EndUser, load_annotation_reply_config
from services.conversation_service import ConversationService
logger = logging.getLogger(__name__)
@@ -37,44 +38,48 @@ class ChatAppGenerator(MessageBasedAppGenerator):
@overload
def generate(
self,
session: Session,
app_model: App,
user: Account | EndUser,
args: Mapping[str, Any],
invoke_from: InvokeFrom,
streaming: Literal[True],
*,
session: Session,
) -> Generator[Mapping | str, None, None]: ...
@overload
def generate(
self,
session: Session,
app_model: App,
user: Account | EndUser,
args: Mapping[str, Any],
invoke_from: InvokeFrom,
streaming: Literal[False],
*,
session: Session,
) -> Mapping[str, Any]: ...
@overload
def generate(
self,
session: Session,
app_model: App,
user: Account | EndUser,
args: Mapping[str, Any],
invoke_from: InvokeFrom,
streaming: bool,
*,
session: Session,
) -> Mapping[str, Any] | Generator[Mapping[str, Any] | str, None, None]: ...
def generate(
self,
session: Session,
app_model: App,
user: Account | EndUser,
args: Mapping[str, Any],
invoke_from: InvokeFrom,
streaming: bool = True,
*,
session: Session,
) -> Mapping[str, Any] | Generator[Mapping[str, Any] | str, None, None]:
"""
Generate App response.
@@ -105,10 +110,14 @@ class ChatAppGenerator(MessageBasedAppGenerator):
conversation_id = args.get("conversation_id")
if conversation_id:
conversation = ConversationService.get_conversation(
app_model=app_model, conversation_id=conversation_id, user=user, session=db.session()
app_model=app_model, conversation_id=conversation_id, user=user, session=session
)
# get app model config
app_model_config = self._get_app_model_config(app_model=app_model, conversation=conversation)
app_model_config = self._get_app_model_config(
app_model=app_model,
conversation=conversation,
session=session,
)
# validate override model config
override_model_config_dict = None
@@ -118,12 +127,16 @@ class ChatAppGenerator(MessageBasedAppGenerator):
# validate config
override_model_config_dict = ChatAppConfigManager.config_validate(
tenant_id=app_model.tenant_id, config=args.get("model_config", {})
tenant_id=app_model.tenant_id, config=args.get("model_config", {}), session=session
)
# always enable retriever resource in debugger mode
override_model_config_dict["retriever_resource"] = {"enabled": True}
annotation_reply = (
None if override_model_config_dict else load_annotation_reply_config(session, app_model_config.app_id)
)
# parse files
# TODO(QuantumGhost): Move file parsing logic to the API controller layer
# for better separation of concerns.
@@ -133,7 +146,7 @@ class ChatAppGenerator(MessageBasedAppGenerator):
with self._bind_file_access_scope(tenant_id=app_model.tenant_id, user=user, invoke_from=invoke_from):
files = args["files"] if args.get("files") else []
file_extra_config = FileUploadConfigManager.convert(
override_model_config_dict or app_model_config.to_dict()
override_model_config_dict or app_model_config.to_dict(annotation_reply=annotation_reply)
)
if file_extra_config:
file_objs = file_factory.build_from_mappings(
@@ -151,6 +164,7 @@ class ChatAppGenerator(MessageBasedAppGenerator):
app_model_config=app_model_config,
conversation=conversation,
override_config_dict=override_model_config_dict,
annotation_reply=annotation_reply,
)
# get tracing instance
@@ -183,7 +197,11 @@ class ChatAppGenerator(MessageBasedAppGenerator):
)
# init generate records
(conversation, message) = self._init_generate_records(application_generate_entity, conversation)
(conversation, message) = self._init_generate_records(
application_generate_entity,
conversation,
session=session,
)
# init queue manager
queue_manager = MessageBasedAppQueueManager(
@@ -202,7 +220,6 @@ class ChatAppGenerator(MessageBasedAppGenerator):
def worker_with_context():
return context.run(
self._generate_worker,
session=session,
flask_app=current_app._get_current_object(), # type: ignore
application_generate_entity=application_generate_entity,
queue_manager=queue_manager,
@@ -229,7 +246,6 @@ class ChatAppGenerator(MessageBasedAppGenerator):
def _generate_worker(
self,
flask_app: Flask,
session: Session,
application_generate_entity: ChatAppGenerateEntity,
queue_manager: AppQueueManager,
conversation_id: str,
@@ -252,13 +268,14 @@ class ChatAppGenerator(MessageBasedAppGenerator):
# chatbot app
runner = ChatAppRunner()
runner.run(
session=session,
application_generate_entity=application_generate_entity,
queue_manager=queue_manager,
conversation=conversation,
message=message,
)
with session_factory.create_session() as session:
runner.run(
application_generate_entity=application_generate_entity,
queue_manager=queue_manager,
conversation=conversation,
message=message,
session=session,
)
except GenerateTaskStoppedError:
pass
except InvokeAuthorizationError:
+15 -11
View File
@@ -17,7 +17,6 @@ from core.memory.token_buffer_memory import TokenBufferMemory
from core.model_manager import ModelInstance
from core.moderation.base import ModerationError
from core.rag.retrieval.dataset_retrieval import DatasetRetrieval
from extensions.ext_database import db
from graphon.file import File
from graphon.model_runtime.entities.message_entities import ImagePromptMessageContent
from models.model import App, Conversation, Message
@@ -32,14 +31,17 @@ class ChatAppRunner(AppRunner):
def run(
self,
session: Session,
application_generate_entity: ChatAppGenerateEntity,
queue_manager: AppQueueManager,
conversation: Conversation,
message: Message,
session: Session,
):
"""
Run application
"""Run the application without retaining ``session`` during model I/O.
Database preparation is committed and the connection is released before
the provider response is requested or consumed.
:param application_generate_entity: application generate entity
:param queue_manager: application queue manager
:param conversation: conversation
@@ -49,10 +51,10 @@ class ChatAppRunner(AppRunner):
app_config = application_generate_entity.app_config
app_config = cast(ChatAppConfig, app_config)
stmt = select(App).where(App.id == app_config.app_id)
with create_session() as session:
app_record = session.scalar(stmt)
with create_session() as read_session:
app_record = read_session.scalar(stmt)
if app_record:
session.expunge(app_record)
read_session.expunge(app_record)
if not app_record:
raise ValueError("App not found")
@@ -123,7 +125,10 @@ class ChatAppRunner(AppRunner):
query=query,
user_id=application_generate_entity.user_id,
invoke_from=application_generate_entity.invoke_from,
session=session,
)
session.commit()
session.close()
if annotation_reply:
queue_manager.publish(
@@ -188,6 +193,9 @@ class ChatAppRunner(AppRunner):
)
context_files = retrieved_files or []
session.commit()
session.close()
# reorganize all inputs and template to prompt messages
# Include: prompt template, inputs, query(optional), files(optional)
# memory(optional), external data, dataset context(optional)
@@ -223,10 +231,6 @@ class ChatAppRunner(AppRunner):
model=application_generate_entity.model_conf.model,
)
# Release the Flask scoped session before LLM streaming so a checked-out DB connection
# is not held for the lifetime of the provider response.
db.session.close()
invoke_result = model_instance.invoke_llm(
prompt_messages=prompt_messages,
model_parameters=application_generate_entity.model_conf.parameters,
@@ -1,5 +1,7 @@
from typing import Any, cast
from sqlalchemy.orm import Session
from core.app.app_config.base_app_config_manager import BaseAppConfigManager
from core.app.app_config.common.sensitive_word_avoidance.manager import SensitiveWordAvoidanceConfigManager
from core.app.app_config.easy_ui_based_app.dataset.manager import DatasetConfigManager
@@ -10,7 +12,7 @@ from core.app.app_config.entities import EasyUIBasedAppConfig, EasyUIBasedAppMod
from core.app.app_config.features.file_upload.manager import FileUploadConfigManager
from core.app.app_config.features.more_like_this.manager import MoreLikeThisConfigManager
from core.app.app_config.features.text_to_speech.manager import TextToSpeechConfigManager
from models.model import App, AppMode, AppModelConfig, AppModelConfigDict
from models.model import AnnotationReplyConfig, App, AppMode, AppModelConfig, AppModelConfigDict
class CompletionAppConfig(EasyUIBasedAppConfig):
@@ -24,7 +26,12 @@ class CompletionAppConfig(EasyUIBasedAppConfig):
class CompletionAppConfigManager(BaseAppConfigManager):
@classmethod
def get_app_config(
cls, app_model: App, app_model_config: AppModelConfig, override_config_dict: AppModelConfigDict | None = None
cls,
app_model: App,
app_model_config: AppModelConfig,
override_config_dict: AppModelConfigDict | None = None,
*,
annotation_reply: AnnotationReplyConfig | None,
) -> CompletionAppConfig:
"""
Convert app model config to completion app config
@@ -39,7 +46,7 @@ class CompletionAppConfigManager(BaseAppConfigManager):
config_from = EasyUIBasedAppModelConfigFrom.APP_LATEST_CONFIG
if config_from != EasyUIBasedAppModelConfigFrom.ARGS:
app_model_config_dict = app_model_config.to_dict()
app_model_config_dict = app_model_config.to_dict(annotation_reply=annotation_reply)
config_dict = app_model_config_dict.copy()
else:
if not override_config_dict:
@@ -68,7 +75,7 @@ class CompletionAppConfigManager(BaseAppConfigManager):
return app_config
@classmethod
def config_validate(cls, tenant_id: str, config: dict[str, Any]) -> AppModelConfigDict:
def config_validate(cls, tenant_id: str, config: dict[str, Any], session: Session) -> AppModelConfigDict:
"""
Validate for completion app model config
@@ -97,7 +104,7 @@ class CompletionAppConfigManager(BaseAppConfigManager):
# dataset_query_variable
config, current_related_config_keys = DatasetConfigManager.validate_and_set_defaults(
tenant_id, app_mode, config
tenant_id, app_mode, config, session
)
related_config_keys.extend(current_related_config_keys)
+61 -28
View File
@@ -21,12 +21,14 @@ from core.app.apps.exc import GenerateTaskStoppedError
from core.app.apps.message_based_app_generator import MessageBasedAppGenerator
from core.app.apps.message_based_app_queue_manager import MessageBasedAppQueueManager
from core.app.entities.app_invoke_entities import CompletionAppGenerateEntity, InvokeFrom
from core.db.session_factory import session_factory
from core.helper.trace_id_helper import extract_trace_session_id_from_args
from core.ops.ops_trace_manager import TraceQueueManager
from extensions.ext_database import db
from factories import file_factory
from graphon.model_runtime.errors.invoke import InvokeAuthorizationError
from models import Account, App, EndUser, Message
from models import Account, App, AppModelConfig, Conversation, EndUser, Message
from models.model import load_annotation_reply_config
from services.errors.app import MoreLikeThisDisabledError
from services.errors.message import MessageNotExistsError
@@ -37,44 +39,48 @@ class CompletionAppGenerator(MessageBasedAppGenerator):
@overload
def generate(
self,
session: Session,
app_model: App,
user: Account | EndUser,
args: Mapping[str, Any],
invoke_from: InvokeFrom,
streaming: Literal[True],
*,
session: Session,
) -> Generator[str | Mapping[str, Any], None, None]: ...
@overload
def generate(
self,
session: Session,
app_model: App,
user: Account | EndUser,
args: Mapping[str, Any],
invoke_from: InvokeFrom,
streaming: Literal[False],
*,
session: Session,
) -> Mapping[str, Any]: ...
@overload
def generate(
self,
session: Session,
app_model: App,
user: Account | EndUser,
args: Mapping[str, Any],
invoke_from: InvokeFrom,
streaming: bool = False,
*,
session: Session,
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]: ...
def generate(
self,
session: Session,
app_model: App,
user: Account | EndUser,
args: Mapping[str, Any],
invoke_from: InvokeFrom,
streaming: bool = True,
*,
session: Session,
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]:
"""
Generate App response.
@@ -96,7 +102,11 @@ class CompletionAppGenerator(MessageBasedAppGenerator):
conversation = None
# get app model config
app_model_config = self._get_app_model_config(app_model=app_model, conversation=conversation)
app_model_config = self._get_app_model_config(
app_model=app_model,
conversation=conversation,
session=session,
)
# validate override model config
override_model_config_dict = None
@@ -106,9 +116,13 @@ class CompletionAppGenerator(MessageBasedAppGenerator):
# validate config
override_model_config_dict = CompletionAppConfigManager.config_validate(
tenant_id=app_model.tenant_id, config=args.get("model_config", {})
tenant_id=app_model.tenant_id, config=args.get("model_config", {}), session=session
)
annotation_reply = (
None if override_model_config_dict else load_annotation_reply_config(session, app_model_config.app_id)
)
# parse files
# TODO(QuantumGhost): Move file parsing logic to the API controller layer
# for better separation of concerns.
@@ -118,7 +132,7 @@ class CompletionAppGenerator(MessageBasedAppGenerator):
with self._bind_file_access_scope(tenant_id=app_model.tenant_id, user=user, invoke_from=invoke_from):
files = args["files"] if args.get("files") else []
file_extra_config = FileUploadConfigManager.convert(
override_model_config_dict or app_model_config.to_dict()
override_model_config_dict or app_model_config.to_dict(annotation_reply=annotation_reply)
)
if file_extra_config:
file_objs = file_factory.build_from_mappings(
@@ -132,7 +146,10 @@ class CompletionAppGenerator(MessageBasedAppGenerator):
# convert to app config
app_config = CompletionAppConfigManager.get_app_config(
app_model=app_model, app_model_config=app_model_config, override_config_dict=override_model_config_dict
app_model=app_model,
app_model_config=app_model_config,
override_config_dict=override_model_config_dict,
annotation_reply=annotation_reply,
)
# get tracing instance
@@ -161,7 +178,10 @@ class CompletionAppGenerator(MessageBasedAppGenerator):
)
# init generate records
(conversation, message) = self._init_generate_records(application_generate_entity)
(conversation, message) = self._init_generate_records(
application_generate_entity,
session=session,
)
# init queue manager
queue_manager = MessageBasedAppQueueManager(
@@ -180,7 +200,6 @@ class CompletionAppGenerator(MessageBasedAppGenerator):
def worker_with_context():
return context.run(
self._generate_worker,
session=session,
flask_app=current_app._get_current_object(), # type: ignore
application_generate_entity=application_generate_entity,
queue_manager=queue_manager,
@@ -206,7 +225,6 @@ class CompletionAppGenerator(MessageBasedAppGenerator):
def _generate_worker(
self,
flask_app: Flask,
session: Session,
application_generate_entity: CompletionAppGenerateEntity,
queue_manager: AppQueueManager,
message_id: str,
@@ -226,12 +244,13 @@ class CompletionAppGenerator(MessageBasedAppGenerator):
# chatbot app
runner = CompletionAppRunner()
runner.run(
session=session,
application_generate_entity=application_generate_entity,
queue_manager=queue_manager,
message=message,
)
with session_factory.create_session() as session:
runner.run(
application_generate_entity=application_generate_entity,
queue_manager=queue_manager,
message=message,
session=session,
)
except GenerateTaskStoppedError:
pass
except InvokeAuthorizationError:
@@ -253,12 +272,13 @@ class CompletionAppGenerator(MessageBasedAppGenerator):
def generate_more_like_this(
self,
session: Session,
app_model: App,
message_id: str,
user: Account | EndUser,
invoke_from: InvokeFrom,
stream: bool = True,
*,
session: Session,
) -> Mapping | Generator[Mapping | str, None, None]:
"""
Generate App response.
@@ -276,12 +296,14 @@ class CompletionAppGenerator(MessageBasedAppGenerator):
Message.from_end_user_id == (user.id if isinstance(user, EndUser) else None),
Message.from_account_id == (user.id if isinstance(user, Account) else None),
)
message = db.session.scalar(stmt)
message = session.scalar(stmt)
if not message:
raise MessageNotExistsError()
current_app_model_config = app_model.app_model_config
current_app_model_config = (
session.get(AppModelConfig, app_model.app_model_config_id) if app_model.app_model_config_id else None
)
if not current_app_model_config:
raise MoreLikeThisDisabledError()
@@ -290,10 +312,16 @@ class CompletionAppGenerator(MessageBasedAppGenerator):
if not current_app_model_config.more_like_this or more_like_this.get("enabled", False) is False:
raise MoreLikeThisDisabledError()
app_model_config = message.app_model_config
conversation = session.get(Conversation, message.conversation_id) if message.conversation_id else None
app_model_config = (
session.get(AppModelConfig, conversation.app_model_config_id)
if conversation and conversation.app_model_config_id
else None
)
if not app_model_config:
raise ValueError("Message app_model_config is None")
override_model_config_dict = app_model_config.to_dict()
annotation_reply = load_annotation_reply_config(session, app_model_config.app_id)
override_model_config_dict = app_model_config.to_dict(annotation_reply=annotation_reply)
model_dict = override_model_config_dict["model"]
completion_params = model_dict.get("completion_params", {})
completion_params["temperature"] = 0.9
@@ -305,7 +333,7 @@ class CompletionAppGenerator(MessageBasedAppGenerator):
file_extra_config = FileUploadConfigManager.convert(override_model_config_dict)
if file_extra_config:
file_objs = file_factory.build_from_mappings(
mappings=message.message_files,
mappings=message.message_files_with_session(session=session),
tenant_id=app_model.tenant_id,
config=file_extra_config,
access_controller=self._file_access_controller,
@@ -315,7 +343,10 @@ class CompletionAppGenerator(MessageBasedAppGenerator):
# convert to app config
app_config = CompletionAppConfigManager.get_app_config(
app_model=app_model, app_model_config=app_model_config, override_config_dict=override_model_config_dict
app_model=app_model,
app_model_config=app_model_config,
override_config_dict=override_model_config_dict,
annotation_reply=annotation_reply,
)
# init application generate entity
@@ -323,7 +354,7 @@ class CompletionAppGenerator(MessageBasedAppGenerator):
task_id=str(uuid.uuid4()),
app_config=app_config,
model_conf=ModelConfigConverter.convert(app_config),
inputs=message.inputs,
inputs=message.inputs_with_session(session=session),
query=message.query,
files=list(file_objs),
user_id=user.id,
@@ -333,7 +364,10 @@ class CompletionAppGenerator(MessageBasedAppGenerator):
)
# init generate records
(conversation, message) = self._init_generate_records(application_generate_entity)
(conversation, message) = self._init_generate_records(
application_generate_entity,
session=session,
)
# init queue manager
queue_manager = MessageBasedAppQueueManager(
@@ -352,7 +386,6 @@ class CompletionAppGenerator(MessageBasedAppGenerator):
def worker_with_context():
return context.run(
self._generate_worker,
session=session,
flask_app=current_app._get_current_object(), # type: ignore
application_generate_entity=application_generate_entity,
queue_manager=queue_manager,
+12 -11
View File
@@ -15,7 +15,6 @@ from core.db.session_factory import create_session
from core.model_manager import ModelInstance
from core.moderation.base import ModerationError
from core.rag.retrieval.dataset_retrieval import DatasetRetrieval
from extensions.ext_database import db
from graphon.file import File
from graphon.model_runtime.entities.message_entities import ImagePromptMessageContent
from models.model import App, Message
@@ -30,13 +29,16 @@ class CompletionAppRunner(AppRunner):
def run(
self,
session: Session,
application_generate_entity: CompletionAppGenerateEntity,
queue_manager: AppQueueManager,
message: Message,
session: Session,
):
"""
Run application
"""Run the application without retaining ``session`` during model I/O.
Database preparation is committed and the connection is released before
the provider response is requested or consumed.
:param application_generate_entity: application generate entity
:param queue_manager: application queue manager
:param message: message
@@ -45,10 +47,10 @@ class CompletionAppRunner(AppRunner):
app_config = application_generate_entity.app_config
app_config = cast(CompletionAppConfig, app_config)
stmt = select(App).where(App.id == app_config.app_id)
with create_session() as session:
app_record = session.scalar(stmt)
with create_session() as read_session:
app_record = read_session.scalar(stmt)
if app_record:
session.expunge(app_record)
read_session.expunge(app_record)
if not app_record:
raise ValueError("App not found")
@@ -150,6 +152,9 @@ class CompletionAppRunner(AppRunner):
)
context_files = retrieved_files or []
session.commit()
session.close()
# reorganize all inputs and template to prompt messages
# Include: prompt template, inputs, query(optional), files(optional)
# memory(optional), external data, dataset context(optional)
@@ -184,10 +189,6 @@ class CompletionAppRunner(AppRunner):
model=application_generate_entity.model_conf.model,
)
# Release the Flask scoped session before LLM streaming so a checked-out DB connection
# is not held for the lifetime of the provider response.
db.session.close()
invoke_result = model_instance.invoke_llm(
prompt_messages=prompt_messages,
model_parameters=application_generate_entity.model_conf.parameters,
@@ -89,12 +89,18 @@ class MessageBasedAppGenerator(BaseAppGenerator):
logger.exception("Failed to handle response, conversation_id: %s", conversation.id)
raise e
def _get_app_model_config(self, app_model: App, conversation: Conversation | None = None) -> AppModelConfig:
def _get_app_model_config(
self,
app_model: App,
conversation: Conversation | None = None,
*,
session: Session,
) -> AppModelConfig:
if conversation:
stmt = select(AppModelConfig).where(
AppModelConfig.id == conversation.app_model_config_id, AppModelConfig.app_id == app_model.id
)
app_model_config = db.session.scalar(stmt)
app_model_config = session.scalar(stmt)
if not app_model_config:
raise AppModelConfigBrokenError()
@@ -102,7 +108,7 @@ class MessageBasedAppGenerator(BaseAppGenerator):
if app_model.app_model_config_id is None:
raise AppModelConfigBrokenError()
app_model_config = app_model.app_model_config
app_model_config = session.get(AppModelConfig, app_model.app_model_config_id)
if not app_model_config:
raise AppModelConfigBrokenError()
@@ -118,6 +124,8 @@ class MessageBasedAppGenerator(BaseAppGenerator):
AdvancedChatAppGenerateEntity,
],
conversation: Conversation | None = None,
*,
session: Session,
) -> tuple[Conversation, Message]:
"""
Initialize generate records
@@ -183,9 +191,9 @@ class MessageBasedAppGenerator(BaseAppGenerator):
from_account_id=account_id,
)
db.session.add(conversation)
db.session.flush()
db.session.refresh(conversation)
session.add(conversation)
session.flush()
session.refresh(conversation)
else:
conversation.updated_at = naive_utc_now()
@@ -216,9 +224,9 @@ class MessageBasedAppGenerator(BaseAppGenerator):
app_mode=app_config.app_mode,
)
db.session.add(message)
db.session.flush()
db.session.refresh(message)
session.add(message)
session.flush()
session.refresh(message)
message_files = []
for file in application_generate_entity.files:
@@ -235,16 +243,16 @@ class MessageBasedAppGenerator(BaseAppGenerator):
message_files.append(message_file)
if message_files:
db.session.add_all(message_files)
session.add_all(message_files)
db.session.commit()
session.commit()
if isinstance(application_generate_entity, ConversationAppGenerateEntity):
application_generate_entity.conversation_id = conversation.id
application_generate_entity.is_new_conversation = created_new_conversation
return conversation, message
except Exception:
db.session.rollback()
session.rollback()
raise
def _get_conversation_introduction(self, application_generate_entity: AppGenerateEntity) -> str:
@@ -64,6 +64,7 @@ class PipelineGenerator(BaseAppGenerator):
def generate(
self,
*,
session: Session,
pipeline: Pipeline,
workflow: Workflow,
user: Account | EndUser,
@@ -79,6 +80,7 @@ class PipelineGenerator(BaseAppGenerator):
def generate(
self,
*,
session: Session,
pipeline: Pipeline,
workflow: Workflow,
user: Account | EndUser,
@@ -94,6 +96,7 @@ class PipelineGenerator(BaseAppGenerator):
def generate(
self,
*,
session: Session,
pipeline: Pipeline,
workflow: Workflow,
user: Account | EndUser,
@@ -108,6 +111,7 @@ class PipelineGenerator(BaseAppGenerator):
def generate(
self,
*,
session: Session,
pipeline: Pipeline,
workflow: Workflow,
user: Account | EndUser,
@@ -120,10 +124,9 @@ class PipelineGenerator(BaseAppGenerator):
) -> Mapping[str, Any] | Generator[Mapping | str, None, None] | None:
# Add null check for dataset
with Session(db.engine, expire_on_commit=False) as session:
dataset = pipeline.retrieve_dataset(session)
if not dataset:
raise ValueError("Pipeline dataset is required")
dataset = pipeline.retrieve_dataset(session)
if not dataset:
raise ValueError("Pipeline dataset is required")
inputs: Mapping[str, Any] = args["inputs"]
start_node_id: str = args["start_node_id"]
datasource_type = DatasourceProviderType(args["datasource_type"])
@@ -157,9 +160,9 @@ class PipelineGenerator(BaseAppGenerator):
batch=batch,
document_form=dataset.chunk_structure,
)
db.session.add(document)
session.add(document)
documents.append(document)
db.session.commit()
session.flush()
# run in child thread
rag_pipeline_invoke_entities = []
@@ -177,8 +180,7 @@ class PipelineGenerator(BaseAppGenerator):
pipeline_id=pipeline.id,
created_by=user.id,
)
db.session.add(document_pipeline_execution_log)
db.session.commit()
session.add(document_pipeline_execution_log)
application_generate_entity = RagPipelineGenerateEntity(
task_id=str(uuid.uuid4()),
app_config=pipeline_config,
@@ -227,6 +229,7 @@ class PipelineGenerator(BaseAppGenerator):
)
if invoke_from == InvokeFrom.DEBUGGER or is_retry:
return self._generate(
session=session,
flask_app=current_app._get_current_object(), # type: ignore
context=contextvars.copy_context(),
pipeline=pipeline,
@@ -253,6 +256,8 @@ class PipelineGenerator(BaseAppGenerator):
)
)
if invoke_from == InvokeFrom.PUBLISHED_PIPELINE and not is_retry:
session.commit()
if rag_pipeline_invoke_entities:
RagPipelineTaskProxy(dataset.tenant_id, user.id, rag_pipeline_invoke_entities).delay()
# return batch, dataset, documents
@@ -282,6 +287,7 @@ class PipelineGenerator(BaseAppGenerator):
def _generate(
self,
*,
session: Session,
flask_app: Flask,
context: contextvars.Context,
pipeline: Pipeline,
@@ -310,7 +316,7 @@ class PipelineGenerator(BaseAppGenerator):
"""
with preserve_flask_contexts(flask_app, context_vars=context):
# init queue manager
workflow = db.session.get(Workflow, workflow_id)
workflow = session.get(Workflow, workflow_id)
if not workflow:
raise ValueError(f"Workflow not found: {workflow_id}")
queue_manager = PipelineQueueManager(
@@ -362,6 +368,8 @@ class PipelineGenerator(BaseAppGenerator):
user: Account | EndUser,
args: Mapping[str, Any],
streaming: bool = True,
*,
session: Session,
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]:
"""
Generate App response.
@@ -372,6 +380,7 @@ class PipelineGenerator(BaseAppGenerator):
:param user: account or end user
:param args: request args
:param streaming: is streamed
:param session: database session supplied by the caller
"""
if not node_id:
raise ValueError("node_id is required")
@@ -384,10 +393,9 @@ class PipelineGenerator(BaseAppGenerator):
pipeline=pipeline, workflow=workflow, start_node_id=args.get("start_node_id", "shared")
)
with Session(db.engine) as session:
dataset = pipeline.retrieve_dataset(session)
if not dataset:
raise ValueError("Pipeline dataset is required")
dataset = pipeline.retrieve_dataset(session)
if not dataset:
raise ValueError("Pipeline dataset is required")
# init application generate entity - use RagPipelineGenerateEntity instead
application_generate_entity = RagPipelineGenerateEntity(
@@ -428,7 +436,7 @@ class PipelineGenerator(BaseAppGenerator):
app_id=application_generate_entity.app_config.app_id,
triggered_from=WorkflowNodeExecutionTriggeredFrom.SINGLE_STEP,
)
draft_var_srv = WorkflowDraftVariableService(db.session())
draft_var_srv = WorkflowDraftVariableService(session)
draft_var_srv.prefill_conversation_variable_default_values(workflow, user_id=user.id)
var_loader = DraftVarLoader(
engine=db.engine,
@@ -438,6 +446,7 @@ class PipelineGenerator(BaseAppGenerator):
)
return self._generate(
session=session,
flask_app=current_app._get_current_object(), # type: ignore
pipeline=pipeline,
workflow_id=workflow.id,
@@ -459,6 +468,8 @@ class PipelineGenerator(BaseAppGenerator):
user: Account | EndUser,
args: Mapping[str, Any],
streaming: bool = True,
*,
session: Session,
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]:
"""
Generate App response.
@@ -469,6 +480,7 @@ class PipelineGenerator(BaseAppGenerator):
:param user: account or end user
:param args: request args
:param streaming: is streamed
:param session: database session supplied by the caller
"""
if not node_id:
raise ValueError("node_id is required")
@@ -476,10 +488,9 @@ class PipelineGenerator(BaseAppGenerator):
if args.get("inputs") is None:
raise ValueError("inputs is required")
with Session(db.engine) as session:
dataset = pipeline.retrieve_dataset(session)
if not dataset:
raise ValueError("Pipeline dataset is required")
dataset = pipeline.retrieve_dataset(session)
if not dataset:
raise ValueError("Pipeline dataset is required")
# convert to app config
pipeline_config = PipelineConfigManager.get_pipeline_config(
@@ -524,7 +535,7 @@ class PipelineGenerator(BaseAppGenerator):
app_id=application_generate_entity.app_config.app_id,
triggered_from=WorkflowNodeExecutionTriggeredFrom.SINGLE_STEP,
)
draft_var_srv = WorkflowDraftVariableService(db.session())
draft_var_srv = WorkflowDraftVariableService(session)
draft_var_srv.prefill_conversation_variable_default_values(workflow, user_id=user.id)
var_loader = DraftVarLoader(
engine=db.engine,
@@ -534,6 +545,7 @@ class PipelineGenerator(BaseAppGenerator):
)
return self._generate(
session=session,
flask_app=current_app._get_current_object(), # type: ignore
pipeline=pipeline,
workflow_id=workflow.id,
+8 -2
View File
@@ -419,6 +419,8 @@ class WorkflowAppGenerator(BaseAppGenerator):
user: Account | EndUser,
args: Mapping[str, Any],
streaming: bool = True,
*,
session: Session,
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]:
"""
Generate App response.
@@ -429,6 +431,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
:param user: account or end user
:param args: request args
:param streaming: is streamed
:param session: database session supplied by the caller
"""
if not node_id:
raise ValueError("node_id is required")
@@ -478,7 +481,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
app_id=application_generate_entity.app_config.app_id,
triggered_from=WorkflowNodeExecutionTriggeredFrom.SINGLE_STEP,
)
draft_var_srv = WorkflowDraftVariableService(db.session())
draft_var_srv = WorkflowDraftVariableService(session)
draft_var_srv.prefill_conversation_variable_default_values(workflow, user_id=user.id)
var_loader = DraftVarLoader(
engine=db.engine,
@@ -508,6 +511,8 @@ class WorkflowAppGenerator(BaseAppGenerator):
user: Account | EndUser,
args: LoopNodeRunPayload,
streaming: bool = True,
*,
session: Session,
) -> Mapping[str, Any] | Generator[str | Mapping[str, Any], None, None]:
"""
Generate App response.
@@ -518,6 +523,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
:param user: account or end user
:param args: request args
:param streaming: is streamed
:param session: database session supplied by the caller
"""
if not node_id:
raise ValueError("node_id is required")
@@ -565,7 +571,7 @@ class WorkflowAppGenerator(BaseAppGenerator):
app_id=application_generate_entity.app_config.app_id,
triggered_from=WorkflowNodeExecutionTriggeredFrom.SINGLE_STEP,
)
draft_var_srv = WorkflowDraftVariableService(db.session())
draft_var_srv = WorkflowDraftVariableService(session)
draft_var_srv.prefill_conversation_variable_default_values(workflow, user_id=user.id)
var_loader = DraftVarLoader(
engine=db.engine,
@@ -1,15 +1,14 @@
import logging
from typing import cast
from sqlalchemy import select
from sqlalchemy.orm import Session
from core.app.entities.app_invoke_entities import InvokeFrom
from core.rag.datasource.vdb.vector_factory import Vector
from core.rag.index_processor.constant.index_type import IndexTechniqueType
from extensions.ext_database import db
from models.dataset import Dataset, DatasetCollectionBinding
from models.dataset import Dataset
from models.enums import CollectionBindingType, ConversationFromSource
from models.model import App, AppAnnotationSetting, Message, MessageAnnotation
from models.model import AnnotationReplyEnabledConfig, App, Message, MessageAnnotation, load_annotation_reply_config
from services.annotation_service import AppAnnotationService
from services.dataset_service import DatasetCollectionBindingService
@@ -25,34 +24,27 @@ class AnnotationReplyFeature:
user_id: str,
invoke_from: InvokeFrom,
*,
session: Session | None = None,
session: Session,
) -> MessageAnnotation | None:
"""Return the closest annotation reply and record a hit in ``session``.
The caller may provide its transaction so the setting lookup, annotation
lookup, and hit-history write share one session. Runtime callers that do
not provide one continue to use Flask-SQLAlchemy's scoped session.
Vector-search failures are logged and return ``None``; transaction
cleanup remains the caller's responsibility.
The setting lookup, vector access, annotation lookup, and hit-history
write share the caller-owned session. Vector-search failures are logged
and return ``None``; transaction cleanup remains the caller's responsibility.
"""
if session is None:
session = db.session()
stmt = select(AppAnnotationSetting).where(AppAnnotationSetting.app_id == app_record.id)
annotation_setting = session.scalar(stmt)
if not annotation_setting:
try:
annotation_reply_config = load_annotation_reply_config(session, app_record.id)
except ValueError:
return None
collection_binding_detail = session.get(DatasetCollectionBinding, annotation_setting.collection_binding_id)
if not collection_binding_detail:
if not annotation_reply_config["enabled"]:
return None
enabled_config = cast(AnnotationReplyEnabledConfig, annotation_reply_config)
try:
score_threshold = annotation_setting.score_threshold or 1
embedding_provider_name = collection_binding_detail.provider_name
embedding_model_name = collection_binding_detail.model_name
score_threshold = enabled_config["score_threshold"] or 1
embedding_provider_name = enabled_config["embedding_model"]["embedding_provider_name"]
embedding_model_name = enabled_config["embedding_model"]["embedding_model_name"]
dataset_collection_binding = DatasetCollectionBindingService.get_dataset_collection_binding(
embedding_provider_name, embedding_model_name, session, CollectionBindingType.ANNOTATION
@@ -67,7 +59,7 @@ class AnnotationReplyFeature:
collection_binding_id=dataset_collection_binding.id,
)
vector = Vector(dataset, attributes=["doc_id", "annotation_id", "app_id"])
vector = Vector(dataset, attributes=["doc_id", "annotation_id", "app_id"], session=session)
documents = vector.search_by_vector(
query=query, top_k=1, score_threshold=score_threshold, filter={"group_id": [dataset.id]}
@@ -5,7 +5,7 @@ from threading import Thread
from typing import Any, cast
from sqlalchemy import select
from sqlalchemy.orm import Session, sessionmaker
from sqlalchemy.orm import Session
from constants.tts_auto_play_timeout import TTS_AUTO_PLAY_TIMEOUT, TTS_AUTO_PLAY_YIELD_CPU_TIME
from core.app.apps.base_app_queue_manager import AppQueueManager, PublishFrom
@@ -44,15 +44,15 @@ from core.app.entities.task_entities import (
)
from core.app.task_pipeline.based_generate_task_pipeline import BasedGenerateTaskPipeline
from core.app.task_pipeline.message_cycle_manager import MessageCycleManager
from core.app.task_pipeline.message_file_utils import MessageFileInfoDict, prepare_file_dict
from core.app.task_pipeline.message_file_utils import prepare_file_dict
from core.base.tts import AppGeneratorTTSPublisher, AudioTrunk
from core.db.session_factory import session_factory
from core.model_manager import ModelInstance
from core.ops.entities.trace_entity import TraceTaskName
from core.ops.ops_trace_manager import TraceQueueManager, TraceTask
from core.prompt.utils.prompt_message_util import PromptMessageUtil
from core.prompt.utils.prompt_template_parser import PromptTemplateParser
from events.message_event import message_was_created
from extensions.ext_database import db
from graphon.file import FileTransferMethod
from graphon.model_runtime.entities.llm_entities import LLMResult, LLMResultChunk, LLMResultChunkDelta, LLMUsage
from graphon.model_runtime.entities.message_entities import (
@@ -269,8 +269,9 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline[EasyUIAppGenerat
match event:
case QueueErrorEvent():
with sessionmaker(bind=db.engine).begin() as session:
with session_factory.create_session() as session:
err = self.handle_error(event=event, session=session, message_id=self._message_id)
session.commit()
yield self.error_to_stream_response(err)
break
case QueueStopEvent() | QueueMessageEndEvent():
@@ -290,17 +291,22 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline[EasyUIAppGenerat
answer=output_moderation_answer
)
with sessionmaker(bind=db.engine).begin() as session:
with session_factory.create_session() as session:
# Save message
self._save_message(session=session, trace_manager=trace_manager)
session.commit()
message_end_resp = self._message_end_to_stream_response()
yield message_end_resp
case QueueRetrieverResourcesEvent():
self._message_cycle_manager.handle_retriever_resources(event)
case QueueAnnotationReplyEvent():
annotation = self._message_cycle_manager.handle_annotation_reply(event)
if annotation:
self._task_state.llm_result.message.content = annotation.content
annotation_content = None
with session_factory.create_session() as session:
annotation = self._message_cycle_manager.handle_annotation_reply(event, session)
if annotation:
annotation_content = annotation.content
if annotation_content:
self._task_state.llm_result.message.content = annotation_content
case QueueAgentThoughtEvent():
agent_thought_response = self._agent_thought_to_stream_response(event)
if agent_thought_response is not None:
@@ -477,8 +483,8 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline[EasyUIAppGenerat
metadata_dict = self._task_state.metadata.model_dump(exclude_none=True)
# Fetch files associated with this message
files: list[MessageFileInfoDict] = []
with Session(db.engine, expire_on_commit=False) as session:
files: Sequence[Mapping[str, Any]] = []
with session_factory.create_session() as session:
message_files = session.scalars(select(MessageFile).where(MessageFile.message_id == self._message_id)).all()
if message_files:
@@ -500,13 +506,13 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline[EasyUIAppGenerat
file_dict = prepare_file_dict(message_file, upload_files_map)
files_list.append(file_dict)
files = files_list
files = cast(Sequence[Mapping[str, Any]], files_list)
return MessageEndStreamResponse(
task_id=self._application_generate_entity.task_id,
id=self._message_id,
metadata=metadata_dict,
files=cast(Sequence[Mapping[str, Any]], files),
files=files,
)
def _agent_message_to_stream_response(self, answer: str, message_id: str) -> AgentMessageStreamResponse:
@@ -526,7 +532,7 @@ class EasyUIBasedGenerateTaskPipeline(BasedGenerateTaskPipeline[EasyUIAppGenerat
:param event: agent thought event
:return:
"""
with Session(db.engine, expire_on_commit=False) as session:
with session_factory.create_session() as session:
agent_thought: MessageAgentThought | None = session.scalar(
select(MessageAgentThought).where(MessageAgentThought.id == event.agent_thought_id).limit(1)
)
@@ -35,7 +35,8 @@ from core.tools.signature import sign_tool_file
from extensions.ext_database import db
from extensions.ext_redis import redis_client
from models.enums import MessageFileBelongsTo
from models.model import AppMode, Conversation, MessageAnnotation, MessageFile
from models.model import App, AppMode, Conversation, MessageAnnotation, MessageFile
from services.account_service import AccountService
from services.annotation_service import AppAnnotationService
logger = logging.getLogger(__name__)
@@ -115,48 +116,50 @@ class MessageCycleManager:
def _generate_conversation_name_worker(self, flask_app: Flask, conversation_id: str, query: str):
with flask_app.app_context():
# get conversation and message
stmt = select(Conversation).where(Conversation.id == conversation_id)
conversation = db.session.scalar(stmt)
with session_factory.create_session() as session:
# get conversation and message
stmt = select(Conversation).where(Conversation.id == conversation_id)
conversation = session.scalar(stmt)
if not conversation:
return
if conversation.mode != AppMode.COMPLETION:
app_model = conversation.app
if not app_model:
if not conversation:
return
# generate conversation name
query_hash = hashlib.md5(query.encode()).hexdigest()[:16]
cache_key = f"conv_name:{conversation_id}:{query_hash}"
if conversation.mode != AppMode.COMPLETION:
app_model = session.get(App, conversation.app_id)
if not app_model:
return
cached_name = redis_client.get(cache_key)
if cached_name:
name = cached_name.decode("utf-8")
else:
try:
name = LLMGenerator.generate_conversation_name(
app_model.tenant_id, query, conversation_id, conversation.app_id
)
redis_client.setex(cache_key, 3600, name)
except Exception:
if dify_config.DEBUG:
logger.exception("generate conversation name failed, conversation_id: %s", conversation_id)
name = query[:47] + "..." if len(query) > 50 else query
conversation.name = name
db.session.commit()
db.session.close()
# generate conversation name
query_hash = hashlib.md5(query.encode()).hexdigest()[:16]
cache_key = f"conv_name:{conversation_id}:{query_hash}"
def handle_annotation_reply(self, event: QueueAnnotationReplyEvent) -> MessageAnnotation | None:
cached_name = redis_client.get(cache_key)
if cached_name:
name = cached_name.decode("utf-8")
else:
try:
name = LLMGenerator.generate_conversation_name(
app_model.tenant_id, query, conversation_id, conversation.app_id
)
redis_client.setex(cache_key, 3600, name)
except Exception:
if dify_config.DEBUG:
logger.exception(
"generate conversation name failed, conversation_id: %s", conversation_id
)
name = query[:47] + "..." if len(query) > 50 else query
conversation.name = name
session.commit()
def handle_annotation_reply(self, event: QueueAnnotationReplyEvent, session: Session) -> MessageAnnotation | None:
"""
Handle annotation reply.
:param event: event
:return:
"""
annotation = AppAnnotationService.get_annotation_by_id(event.message_annotation_id, session=db.session())
annotation = AppAnnotationService.get_annotation_by_id(event.message_annotation_id, session)
if annotation:
account = annotation.account
account = AccountService.get_account_by_id(annotation.account_id, session=session)
self._task_state.metadata.annotation_reply = AnnotationReply(
id=annotation.id,
account=AnnotationReplyAccount(
+177 -121
View File
@@ -10,9 +10,11 @@ from typing import Any
from flask import Flask, current_app
from sqlalchemy import delete, func, select, update
from sqlalchemy.orm import Session
from sqlalchemy.orm.exc import ObjectDeletedError
from configs import dify_config
from core.db.session_factory import session_factory
from core.entities.knowledge_entities import IndexingEstimate, PreviewDetail, QAPreviewDetail
from core.errors.error import ProviderTokenNotInitError
from core.model_manager import ModelInstance, ModelManager
@@ -31,7 +33,6 @@ from core.rag.splitter.fixed_text_splitter import (
)
from core.rag.splitter.text_splitter import TextSplitter
from core.tools.utils.web_reader_tool import get_image_upload_file_ids
from extensions.ext_database import db
from extensions.ext_redis import redis_client
from extensions.ext_storage import storage
from graphon.model_runtime.entities.model_entities import ModelType
@@ -54,30 +55,34 @@ class IndexingRunner:
def _get_model_manager(tenant_id: str) -> ModelManager:
return ModelManager.for_tenant(tenant_id=tenant_id)
def _handle_indexing_error(self, document_id: str, error: Exception) -> None:
def _handle_indexing_error(self, document_id: str, error: Exception, session: Session) -> None:
"""Handle indexing errors by updating document status."""
logger.exception("consume document failed")
document = db.session.get(DatasetDocument, document_id)
document = session.get(DatasetDocument, document_id)
if document:
document.indexing_status = IndexingStatus.ERROR
error_message = getattr(error, "description", str(error))
document.error = str(error_message)
document.stopped_at = naive_utc_now()
db.session.commit()
session.flush()
def run(self, dataset_documents: list[DatasetDocument]):
"""Run the indexing process."""
def run(self, dataset_documents: list[DatasetDocument], session: Session):
"""Run indexing with commits before slow transforms and parallel index workers.
The phase commits keep document locks short and make newly created segments
visible to the worker sessions used for keyword and vector indexing.
"""
for dataset_document in dataset_documents:
document_id = dataset_document.id
try:
# Re-query the document to ensure it's bound to the current session
requeried_document = db.session.get(DatasetDocument, document_id)
requeried_document = session.get(DatasetDocument, document_id)
if not requeried_document:
logger.warning("Document not found, skipping document id: %s", document_id)
continue
# get dataset
dataset = db.session.get(Dataset, requeried_document.dataset_id)
dataset = session.get(Dataset, requeried_document.dataset_id)
if not dataset:
raise ValueError("no dataset found")
@@ -85,19 +90,20 @@ class IndexingRunner:
stmt = select(DatasetProcessRule).where(
DatasetProcessRule.id == requeried_document.dataset_process_rule_id
)
processing_rule = db.session.scalar(stmt)
processing_rule = session.scalar(stmt)
if not processing_rule:
raise ValueError("no process rule found")
index_type = requeried_document.doc_form
index_processor = IndexProcessorFactory(index_type).init_index_processor()
# extract
text_docs = self._extract(index_processor, requeried_document, processing_rule.to_dict())
text_docs = self._extract(index_processor, requeried_document, processing_rule.to_dict(), session)
session.commit()
# transform
current_user = db.session.get(Account, requeried_document.created_by)
current_user = session.get(Account, requeried_document.created_by)
if not current_user:
raise ValueError("no current user found")
current_user.set_tenant_id(dataset.tenant_id)
current_user.set_tenant_id_with_session(dataset.tenant_id, session=session)
documents = self._transform(
index_processor,
dataset,
@@ -105,9 +111,11 @@ class IndexingRunner:
requeried_document.doc_language,
processing_rule.to_dict(),
current_user=current_user,
session=session,
)
# save segment
self._load_segments(dataset, requeried_document, documents)
self._load_segments(dataset, requeried_document, documents, session)
session.commit()
# load
self._load(
@@ -115,34 +123,35 @@ class IndexingRunner:
dataset=dataset,
dataset_document=requeried_document,
documents=documents,
session=session,
)
except DocumentIsPausedError:
raise DocumentIsPausedError(f"Document paused, document id: {document_id}")
except ProviderTokenNotInitError as e:
self._handle_indexing_error(document_id, e)
self._handle_indexing_error(document_id, e, session)
except ObjectDeletedError:
logger.warning("Document deleted, document id: %s", document_id)
except Exception as e:
self._handle_indexing_error(document_id, e)
self._handle_indexing_error(document_id, e, session)
def run_in_splitting_status(self, dataset_document: DatasetDocument):
def run_in_splitting_status(self, dataset_document: DatasetDocument, session: Session):
"""Run the indexing process when the index_status is splitting."""
document_id = dataset_document.id
try:
# Re-query the document to ensure it's bound to the current session
requeried_document = db.session.get(DatasetDocument, document_id)
requeried_document = session.get(DatasetDocument, document_id)
if not requeried_document:
logger.warning("Document not found: %s", document_id)
return
# get dataset
dataset = db.session.get(Dataset, requeried_document.dataset_id)
dataset = session.get(Dataset, requeried_document.dataset_id)
if not dataset:
raise ValueError("no dataset found")
# get exist document_segment list and delete
document_segments = db.session.scalars(
document_segments = session.scalars(
select(DocumentSegment).where(
DocumentSegment.dataset_id == dataset.id,
DocumentSegment.document_id == requeried_document.id,
@@ -150,27 +159,28 @@ class IndexingRunner:
).all()
for document_segment in document_segments:
db.session.delete(document_segment)
session.delete(document_segment)
if requeried_document.doc_form == IndexStructureType.PARENT_CHILD_INDEX:
# delete child chunks
db.session.execute(delete(ChildChunk).where(ChildChunk.segment_id == document_segment.id))
db.session.commit()
session.execute(delete(ChildChunk).where(ChildChunk.segment_id == document_segment.id))
session.commit()
# get the process rule
stmt = select(DatasetProcessRule).where(DatasetProcessRule.id == requeried_document.dataset_process_rule_id)
processing_rule = db.session.scalar(stmt)
processing_rule = session.scalar(stmt)
if not processing_rule:
raise ValueError("no process rule found")
index_type = requeried_document.doc_form
index_processor = IndexProcessorFactory(index_type).init_index_processor()
# extract
text_docs = self._extract(index_processor, requeried_document, processing_rule.to_dict())
text_docs = self._extract(index_processor, requeried_document, processing_rule.to_dict(), session)
session.commit()
# transform
current_user = db.session.get(Account, requeried_document.created_by)
current_user = session.get(Account, requeried_document.created_by)
if not current_user:
raise ValueError("no current user found")
current_user.set_tenant_id(dataset.tenant_id)
current_user.set_tenant_id_with_session(dataset.tenant_id, session=session)
documents = self._transform(
index_processor,
dataset,
@@ -178,9 +188,11 @@ class IndexingRunner:
requeried_document.doc_language,
processing_rule.to_dict(),
current_user=current_user,
session=session,
)
# save segment
self._load_segments(dataset, requeried_document, documents)
self._load_segments(dataset, requeried_document, documents, session)
session.commit()
# load
self._load(
@@ -188,32 +200,33 @@ class IndexingRunner:
dataset=dataset,
dataset_document=requeried_document,
documents=documents,
session=session,
)
except DocumentIsPausedError:
raise DocumentIsPausedError(f"Document paused, document id: {document_id}")
except ProviderTokenNotInitError as e:
self._handle_indexing_error(document_id, e)
self._handle_indexing_error(document_id, e, session)
except Exception as e:
self._handle_indexing_error(document_id, e)
self._handle_indexing_error(document_id, e, session)
def run_in_indexing_status(self, dataset_document: DatasetDocument):
def run_in_indexing_status(self, dataset_document: DatasetDocument, session: Session):
"""Run the indexing process when the index_status is indexing."""
document_id = dataset_document.id
try:
# Re-query the document to ensure it's bound to the current session
requeried_document = db.session.get(DatasetDocument, document_id)
requeried_document = session.get(DatasetDocument, document_id)
if not requeried_document:
logger.warning("Document not found: %s", document_id)
return
# get dataset
dataset = db.session.get(Dataset, requeried_document.dataset_id)
dataset = session.get(Dataset, requeried_document.dataset_id)
if not dataset:
raise ValueError("no dataset found")
# get exist document_segment list and delete
document_segments = db.session.scalars(
document_segments = session.scalars(
select(DocumentSegment).where(
DocumentSegment.dataset_id == dataset.id,
DocumentSegment.document_id == requeried_document.id,
@@ -235,7 +248,7 @@ class IndexingRunner:
},
)
if requeried_document.doc_form == IndexStructureType.PARENT_CHILD_INDEX:
child_chunks = document_segment.get_child_chunks()
child_chunks = document_segment.get_child_chunks(session=session)
if child_chunks:
child_documents = []
for child_chunk in child_chunks:
@@ -259,13 +272,14 @@ class IndexingRunner:
dataset=dataset,
dataset_document=requeried_document,
documents=documents,
session=session,
)
except DocumentIsPausedError:
raise DocumentIsPausedError(f"Document paused, document id: {document_id}")
except ProviderTokenNotInitError as e:
self._handle_indexing_error(document_id, e)
self._handle_indexing_error(document_id, e, session)
except Exception as e:
self._handle_indexing_error(document_id, e)
self._handle_indexing_error(document_id, e, session)
def indexing_estimate(
self,
@@ -276,6 +290,8 @@ class IndexingRunner:
doc_language: str = "English",
dataset_id: str | None = None,
indexing_technique: str = IndexTechniqueType.ECONOMY,
*,
session: Session,
) -> IndexingEstimate:
"""
Estimate the indexing for the document.
@@ -289,7 +305,7 @@ class IndexingRunner:
embedding_model_instance = None
if dataset_id:
dataset = db.session.get(Dataset, dataset_id)
dataset = session.get(Dataset, dataset_id)
if not dataset:
raise ValueError("Dataset not found.")
if IndexTechniqueType.HIGH_QUALITY in {dataset.indexing_technique, indexing_technique}:
@@ -316,7 +332,6 @@ class IndexingRunner:
qa_preview_texts: list[QAPreviewDetail] = []
total_segments = 0
deleted_preview_images = False
# doc_form represents the segmentation method (general, parent-child, QA)
index_type = doc_form
index_processor = IndexProcessorFactory(index_type).init_index_processor()
@@ -328,7 +343,9 @@ class IndexingRunner:
"rules": tmp_processing_rule.get("rules"),
}
# Extract document content
text_docs = index_processor.extract(extract_setting, process_rule_mode=tmp_processing_rule["mode"])
text_docs = index_processor.extract(
extract_setting, process_rule_mode=tmp_processing_rule["mode"], session=session
)
# Cleaning and segmentation
documents = index_processor.transform(
text_docs,
@@ -338,6 +355,7 @@ class IndexingRunner:
tenant_id=tenant_id,
doc_language=doc_language,
preview=True,
session=session,
)
total_segments += len(documents)
for document in documents:
@@ -357,7 +375,7 @@ class IndexingRunner:
image_upload_file_ids = get_image_upload_file_ids(document.page_content)
for upload_file_id in image_upload_file_ids:
stmt = select(UploadFile).where(UploadFile.id == upload_file_id)
image_file = db.session.scalar(stmt)
image_file = session.scalar(stmt)
if image_file is None:
continue
try:
@@ -368,11 +386,11 @@ class IndexingRunner:
image_upload_file_is: %s",
upload_file_id,
)
db.session.delete(image_file)
deleted_preview_images = True
session.delete(image_file)
if deleted_preview_images:
db.session.commit()
# Persist preview cleanup and release the caller transaction before
# summary workers query through their own sessions.
session.commit()
if doc_form and doc_form == "qa_model":
return IndexingEstimate(total_segments=total_segments * 20, qa_preview=qa_preview_texts, preview=[])
@@ -381,13 +399,17 @@ class IndexingRunner:
summary_index_setting = tmp_processing_rule.get("summary_index_setting")
if summary_index_setting and summary_index_setting.get("enable") and preview_texts:
preview_texts = index_processor.generate_summary_preview(
tenant_id, preview_texts, summary_index_setting, doc_language
tenant_id, preview_texts, summary_index_setting, doc_language, session=session
)
return IndexingEstimate(total_segments=total_segments, preview=preview_texts)
def _extract(
self, index_processor: BaseIndexProcessor, dataset_document: DatasetDocument, process_rule: Mapping[str, Any]
self,
index_processor: BaseIndexProcessor,
dataset_document: DatasetDocument,
process_rule: Mapping[str, Any],
session: Session,
) -> list[Document]:
data_source_info = dataset_document.data_source_info_dict
text_docs = []
@@ -396,7 +418,7 @@ class IndexingRunner:
if not data_source_info or "upload_file_id" not in data_source_info:
raise ValueError("no upload file found")
stmt = select(UploadFile).where(UploadFile.id == data_source_info["upload_file_id"])
file_detail = db.session.scalars(stmt).one_or_none()
file_detail = session.scalars(stmt).one_or_none()
if file_detail:
extract_setting = ExtractSetting(
@@ -404,7 +426,9 @@ class IndexingRunner:
upload_file=file_detail,
document_model=dataset_document.doc_form,
)
text_docs = index_processor.extract(extract_setting, process_rule_mode=process_rule["mode"])
text_docs = index_processor.extract(
extract_setting, process_rule_mode=process_rule["mode"], session=session
)
case DataSourceType.NOTION_IMPORT:
if (
not data_source_info
@@ -426,7 +450,9 @@ class IndexingRunner:
),
document_model=dataset_document.doc_form,
)
text_docs = index_processor.extract(extract_setting, process_rule_mode=process_rule["mode"])
text_docs = index_processor.extract(
extract_setting, process_rule_mode=process_rule["mode"], session=session
)
case DataSourceType.WEBSITE_CRAWL:
if (
not data_source_info
@@ -449,11 +475,14 @@ class IndexingRunner:
),
document_model=dataset_document.doc_form,
)
text_docs = index_processor.extract(extract_setting, process_rule_mode=process_rule["mode"])
text_docs = index_processor.extract(
extract_setting, process_rule_mode=process_rule["mode"], session=session
)
case _:
return []
# update document status to splitting
self._update_document_index_status(
session=session,
document_id=dataset_document.id,
after_indexing_status=IndexingStatus.SPLITTING,
extra_update_params={
@@ -576,6 +605,7 @@ class IndexingRunner:
dataset: Dataset,
dataset_document: DatasetDocument,
documents: list[Document],
session: Session,
):
"""
insert index and update document/segment status to completed
@@ -625,10 +655,10 @@ class IndexingRunner:
executor.submit(
self._process_chunk,
current_app._get_current_object(), # type: ignore
index_processor,
dataset_document.doc_form,
chunk_documents,
dataset,
dataset_document,
dataset.id,
dataset_document.id,
embedding_model_instance,
)
)
@@ -645,6 +675,7 @@ class IndexingRunner:
# update document status to completed
self._update_document_index_status(
session=session,
document_id=dataset_document.id,
after_indexing_status=IndexingStatus.COMPLETED,
extra_update_params={
@@ -656,20 +687,80 @@ class IndexingRunner:
)
@staticmethod
def _process_keyword_index(flask_app, dataset_id, document_id, documents):
def _process_keyword_index(flask_app: Flask, dataset_id: str, document_id: str, documents: list[Document]):
with flask_app.app_context():
dataset = db.session.get(Dataset, dataset_id)
if not dataset:
raise ValueError("no dataset found")
keyword = Keyword(dataset)
keyword.create(documents)
if dataset.indexing_technique != IndexTechniqueType.HIGH_QUALITY:
document_ids = [document.metadata["doc_id"] for document in documents]
db.session.execute(
with session_factory.create_session() as session:
dataset = session.get(Dataset, dataset_id)
if not dataset:
raise ValueError("no dataset found")
keyword = Keyword(dataset)
keyword.create(documents, session)
if dataset.indexing_technique != IndexTechniqueType.HIGH_QUALITY:
document_ids = [document.metadata["doc_id"] for document in documents]
session.execute(
update(DocumentSegment)
.where(
DocumentSegment.document_id == document_id,
DocumentSegment.dataset_id == dataset_id,
DocumentSegment.index_node_id.in_(document_ids),
DocumentSegment.status == SegmentStatus.INDEXING,
)
.values(
status=SegmentStatus.COMPLETED,
enabled=True,
completed_at=naive_utc_now(),
)
)
session.commit()
def _process_chunk(
self,
flask_app: Flask,
index_type: str,
chunk_documents: list[Document],
dataset_id: str,
dataset_document_id: str,
embedding_model_instance: ModelInstance | None,
):
with flask_app.app_context():
with session_factory.create_session() as session:
dataset = session.get(Dataset, dataset_id)
if not dataset:
raise ValueError("no dataset found")
dataset_document = session.get(DatasetDocument, dataset_document_id)
if not dataset_document:
raise ValueError("no document found")
# check document is paused
self._check_document_paused_status(dataset_document.id)
tokens = 0
if embedding_model_instance:
page_content_list = [document.page_content for document in chunk_documents]
tokens += sum(embedding_model_instance.get_text_embedding_num_tokens(page_content_list))
multimodal_documents = []
for document in chunk_documents:
if document.attachments and dataset.is_multimodal:
multimodal_documents.extend(document.attachments)
# load index
index_processor = IndexProcessorFactory(index_type).init_index_processor()
index_processor.load(
dataset,
chunk_documents,
multimodal_documents=multimodal_documents,
with_keywords=False,
session=session,
)
document_ids = [document.metadata["doc_id"] for document in chunk_documents]
session.execute(
update(DocumentSegment)
.where(
DocumentSegment.document_id == document_id,
DocumentSegment.dataset_id == dataset_id,
DocumentSegment.document_id == dataset_document.id,
DocumentSegment.dataset_id == dataset.id,
DocumentSegment.index_node_id.in_(document_ids),
DocumentSegment.status == SegmentStatus.INDEXING,
)
@@ -680,55 +771,9 @@ class IndexingRunner:
)
)
db.session.commit()
session.commit()
def _process_chunk(
self,
flask_app: Flask,
index_processor: BaseIndexProcessor,
chunk_documents: list[Document],
dataset: Dataset,
dataset_document: DatasetDocument,
embedding_model_instance: ModelInstance | None,
):
with flask_app.app_context():
# check document is paused
self._check_document_paused_status(dataset_document.id)
tokens = 0
if embedding_model_instance:
page_content_list = [document.page_content for document in chunk_documents]
tokens += sum(embedding_model_instance.get_text_embedding_num_tokens(page_content_list))
multimodal_documents = []
for document in chunk_documents:
if document.attachments and dataset.is_multimodal:
multimodal_documents.extend(document.attachments)
# load index
index_processor.load(
dataset, chunk_documents, multimodal_documents=multimodal_documents, with_keywords=False
)
document_ids = [document.metadata["doc_id"] for document in chunk_documents]
db.session.execute(
update(DocumentSegment)
.where(
DocumentSegment.document_id == dataset_document.id,
DocumentSegment.dataset_id == dataset.id,
DocumentSegment.index_node_id.in_(document_ids),
DocumentSegment.status == SegmentStatus.INDEXING,
)
.values(
status=SegmentStatus.COMPLETED,
enabled=True,
completed_at=naive_utc_now(),
)
)
db.session.commit()
return tokens
return tokens
@staticmethod
def _check_document_paused_status(document_id: str):
@@ -742,12 +787,14 @@ class IndexingRunner:
document_id: str,
after_indexing_status: IndexingStatus,
extra_update_params: Mapping[Any, Any] | None = None,
*,
session: Session,
):
"""
Update the document indexing status.
"""
count = (
db.session.scalar(
session.scalar(
select(func.count())
.select_from(DatasetDocument)
.where(DatasetDocument.id == document_id, DatasetDocument.is_paused == True)
@@ -756,7 +803,7 @@ class IndexingRunner:
)
if count > 0:
raise DocumentIsPausedError()
document = db.session.get(DatasetDocument, document_id)
document = session.get(DatasetDocument, document_id)
if not document:
raise DocumentIsDeletedPausedError()
@@ -764,18 +811,18 @@ class IndexingRunner:
if extra_update_params:
update_params.update(extra_update_params)
db.session.execute(update(DatasetDocument).where(DatasetDocument.id == document_id).values(update_params)) # type: ignore
db.session.commit()
session.execute(update(DatasetDocument).where(DatasetDocument.id == document_id).values(update_params)) # type: ignore
session.flush()
@staticmethod
def _update_segments_by_document(dataset_document_id: str, update_params: Mapping[Any, Any]):
def _update_segments_by_document(dataset_document_id: str, update_params: Mapping[Any, Any], session: Session):
"""
Update the document segment by document id.
"""
db.session.execute(
session.execute(
update(DocumentSegment).where(DocumentSegment.document_id == dataset_document_id).values(update_params)
)
db.session.commit()
session.flush()
def _transform(
self,
@@ -785,6 +832,8 @@ class IndexingRunner:
doc_language: str,
process_rule: Mapping[str, Any],
current_user: Account | None = None,
*,
session: Session,
) -> list[Document]:
# get embedding model instance
embedding_model_instance = None
@@ -809,11 +858,14 @@ class IndexingRunner:
process_rule=process_rule,
tenant_id=dataset.tenant_id,
doc_language=doc_language,
session=session,
)
return documents
def _load_segments(self, dataset: Dataset, dataset_document: DatasetDocument, documents: list[Document]):
def _load_segments(
self, dataset: Dataset, dataset_document: DatasetDocument, documents: list[Document], session: Session
):
# save node to document segment
doc_store = DatasetDocumentStore(
dataset=dataset, user_id=dataset_document.created_by, document_id=dataset_document.id
@@ -821,12 +873,15 @@ class IndexingRunner:
# add document segments
doc_store.add_documents(
docs=documents, save_child=dataset_document.doc_form == IndexStructureType.PARENT_CHILD_INDEX
docs=documents,
save_child=dataset_document.doc_form == IndexStructureType.PARENT_CHILD_INDEX,
session=session,
)
# update document status to indexing
cur_time = naive_utc_now()
self._update_document_index_status(
session=session,
document_id=dataset_document.id,
after_indexing_status=IndexingStatus.INDEXING,
extra_update_params={
@@ -838,6 +893,7 @@ class IndexingRunner:
# update segment status to indexing
self._update_segments_by_document(
session=session,
dataset_document_id=dataset_document.id,
update_params={
DocumentSegment.status: SegmentStatus.INDEXING,
+1 -1
View File
@@ -63,6 +63,6 @@ class BaseTraceInstance(ABC):
)
if not current_tenant:
raise ValueError(f"Current tenant not found for account {service_account.id}")
service_account.set_tenant_id(current_tenant.tenant_id)
service_account.set_tenant_id_with_session(current_tenant.tenant_id, session=session)
return service_account
+9 -7
View File
@@ -61,7 +61,6 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation):
@classmethod
def invoke_app(
cls,
session: Session,
app_id: str,
user_id: str,
tenant_id: str,
@@ -70,6 +69,7 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation):
stream: bool,
inputs: Mapping,
files: list[dict],
session: Session,
) -> Generator[Mapping | str, None, None] | Mapping:
"""
invoke app
@@ -91,21 +91,20 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation):
if not query:
raise ValueError("missing query")
return cls.invoke_chat_app(session, app, user, conversation_id, query, stream, inputs, files)
return cls.invoke_chat_app(app, user, conversation_id, query, stream, inputs, files, session)
case AppMode.WORKFLOW:
workflow = cls._get_workflow(app)
if not workflow:
raise ValueError("unexpected app type")
return cls.invoke_workflow_app(app, workflow, user, stream, inputs, files)
case AppMode.COMPLETION:
return cls.invoke_completion_app(session, app, user, stream, inputs, files)
return cls.invoke_completion_app(app, user, stream, inputs, files, session)
case _:
raise ValueError("unexpected app type")
@classmethod
def invoke_chat_app(
cls,
session: Session,
app: App,
user: Account | EndUser,
conversation_id: str,
@@ -113,6 +112,7 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation):
stream: bool,
inputs: Mapping,
files: list[dict],
session: Session,
) -> Generator[Mapping | str, None, None] | Mapping:
"""
invoke chat app
@@ -142,6 +142,7 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation):
workflow_run_id=str(uuid.uuid4()),
streaming=stream,
pause_state_config=pause_config,
session=session,
)
case AppMode.AGENT_CHAT:
return AgentChatAppGenerator().generate(
@@ -155,10 +156,10 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation):
},
invoke_from=InvokeFrom.SERVICE_API,
streaming=stream,
session=session,
)
case AppMode.CHAT:
return ChatAppGenerator().generate(
session=session,
app_model=app,
user=user,
args={
@@ -169,6 +170,7 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation):
},
invoke_from=InvokeFrom.SERVICE_API,
streaming=stream,
session=session,
)
case _:
raise ValueError("unexpected app type")
@@ -205,23 +207,23 @@ class PluginAppBackwardsInvocation(BaseBackwardsInvocation):
@classmethod
def invoke_completion_app(
cls,
session: Session,
app: App,
user: EndUser | Account,
stream: bool,
inputs: Mapping,
files: list[dict],
session: Session,
) -> Generator[Mapping | str, None, None] | Mapping:
"""
invoke completion app
"""
return CompletionAppGenerator().generate(
session=session,
app_model=app,
user=user,
args={"inputs": inputs, "files": files},
invoke_from=InvokeFrom.SERVICE_API,
streaming=stream,
session=session,
)
@classmethod
@@ -1,18 +1,18 @@
from sqlalchemy import select
from sqlalchemy.orm import Session
from core.prompt.utils.extract_thread_messages import extract_thread_messages
from extensions.ext_database import db
from models.model import Message
def get_thread_messages_length(conversation_id: str) -> int:
def get_thread_messages_length(conversation_id: str, *, session: Session) -> int:
"""
Get the number of thread messages based on the parent message id.
"""
# Fetch all messages related to the conversation
stmt = select(Message).where(Message.conversation_id == conversation_id).order_by(Message.created_at.desc())
messages = db.session.scalars(stmt).all()
messages = session.scalars(stmt).all()
# Extract thread messages
thread_messages = extract_thread_messages(messages)
@@ -1,5 +1,7 @@
from typing import TypedDict
from sqlalchemy.orm import Session
from core.model_manager import ModelInstance, ModelManager
from core.rag.data_post_processor.reorder import ReorderRunner
from core.rag.index_processor.constant.query_type import QueryType
@@ -42,8 +44,12 @@ class DataPostProcessor:
reranking_model: RerankingModelDict | None = None,
weights: WeightsDict | None = None,
reorder_enabled: bool = False,
*,
session: Session,
):
self.rerank_runner = self._get_rerank_runner(reranking_mode, tenant_id, reranking_model, weights)
self.rerank_runner = self._get_rerank_runner(
reranking_mode, tenant_id, reranking_model, weights, session=session
)
self.reorder_runner = self._get_reorder_runner(reorder_enabled)
def invoke(
@@ -68,6 +74,8 @@ class DataPostProcessor:
tenant_id: str,
reranking_model: RerankingModelDict | None = None,
weights: WeightsDict | None = None,
*,
session: Session,
) -> BaseRerankRunner | None:
if reranking_mode == RerankMode.WEIGHTED_SCORE and weights:
runner = RerankRunnerFactory.create_rerank_runner(
@@ -90,7 +98,9 @@ class DataPostProcessor:
if rerank_model_instance is None:
return None
runner = RerankRunnerFactory.create_rerank_runner(
runner_type=reranking_mode, rerank_model_instance=rerank_model_instance
runner_type=reranking_mode,
rerank_model_instance=rerank_model_instance,
session=session,
)
return runner
return None
+61 -43
View File
@@ -4,12 +4,12 @@ from typing import Any, TypedDict, override
import orjson
from pydantic import BaseModel
from sqlalchemy import select
from sqlalchemy.orm import Session
from configs import dify_config
from core.rag.datasource.keyword.jieba.jieba_keyword_table_handler import JiebaKeywordTableHandler
from core.rag.datasource.keyword.keyword_base import BaseKeyword
from core.rag.models.document import Document
from extensions.ext_database import db
from extensions.ext_redis import redis_client
from extensions.ext_storage import storage
from models.dataset import Dataset, DatasetKeywordTable, DocumentSegment
@@ -30,32 +30,32 @@ class Jieba(BaseKeyword):
self._config = KeywordTableConfig()
@override
def create(self, texts: list[Document], **kwargs) -> BaseKeyword:
def create(self, texts: list[Document], session: Session, **kwargs: Any) -> BaseKeyword:
lock_name = f"keyword_indexing_lock_{self.dataset.id}"
with redis_client.lock(lock_name, timeout=600):
keyword_table_handler = JiebaKeywordTableHandler()
keyword_table = self._get_dataset_keyword_table()
keyword_table = self._get_dataset_keyword_table(session=session)
keyword_number = self.dataset.keyword_number or self._config.max_keywords_per_chunk
for text in texts:
keywords = keyword_table_handler.extract_keywords(text.page_content, keyword_number)
if text.metadata is not None:
self._update_segment_keywords(self.dataset.id, text.metadata["doc_id"], list(keywords))
self._update_segment_keywords(self.dataset.id, text.metadata["doc_id"], list(keywords), session)
keyword_table = self._add_text_to_keyword_table(
keyword_table or {}, text.metadata["doc_id"], list(keywords)
)
self._save_dataset_keyword_table(keyword_table)
self._save_dataset_keyword_table(keyword_table, session)
return self
@override
def add_texts(self, texts: list[Document], **kwargs):
def add_texts(self, texts: list[Document], session: Session, **kwargs: Any):
lock_name = f"keyword_indexing_lock_{self.dataset.id}"
with redis_client.lock(lock_name, timeout=600):
keyword_table_handler = JiebaKeywordTableHandler()
keyword_table = self._get_dataset_keyword_table()
keyword_table = self._get_dataset_keyword_table(session=session)
keywords_list = kwargs.get("keywords_list")
keyword_number = self.dataset.keyword_number or self._config.max_keywords_per_chunk
for i in range(len(texts)):
@@ -67,33 +67,47 @@ class Jieba(BaseKeyword):
else:
keywords = keyword_table_handler.extract_keywords(text.page_content, keyword_number)
if text.metadata is not None:
self._update_segment_keywords(self.dataset.id, text.metadata["doc_id"], list(keywords))
self._update_segment_keywords(self.dataset.id, text.metadata["doc_id"], list(keywords), session)
keyword_table = self._add_text_to_keyword_table(
keyword_table or {}, text.metadata["doc_id"], list(keywords)
)
self._save_dataset_keyword_table(keyword_table)
self._save_dataset_keyword_table(keyword_table, session)
@override
def text_exists(self, id: str) -> bool:
keyword_table = self._get_dataset_keyword_table()
def text_exists(self, id: str, *, session: Session) -> bool:
dataset_keyword_table = self.dataset.get_dataset_keyword_table(session=session)
keyword_table = None
keyword_table_dict = (
dataset_keyword_table.get_keyword_table_dict(session=session) if dataset_keyword_table else None
)
if keyword_table_dict:
data: Any = keyword_table_dict["__data__"]
keyword_table = dict(data["table"])
if keyword_table is None:
return False
return id in set.union(*keyword_table.values())
@override
def delete_by_ids(self, ids: list[str]):
def delete_by_ids(self, ids: list[str], session: Session, **kwargs: Any):
lock_name = f"keyword_indexing_lock_{self.dataset.id}"
with redis_client.lock(lock_name, timeout=600):
keyword_table = self._get_dataset_keyword_table()
keyword_table = self._get_dataset_keyword_table(session)
if keyword_table is not None:
keyword_table = self._delete_ids_from_keyword_table(keyword_table, ids)
self._save_dataset_keyword_table(keyword_table)
self._save_dataset_keyword_table(keyword_table, session)
@override
def search(self, query: str, **kwargs: Any) -> list[Document]:
keyword_table = self._get_dataset_keyword_table()
def search(self, query: str, *, session: Session, **kwargs: Any) -> list[Document]:
dataset_keyword_table = self.dataset.get_dataset_keyword_table(session=session)
keyword_table = None
keyword_table_dict = (
dataset_keyword_table.get_keyword_table_dict(session=session) if dataset_keyword_table else None
)
if keyword_table_dict:
data: Any = keyword_table_dict["__data__"]
keyword_table = dict(data["table"])
k = kwargs.get("top_k", 4)
document_ids_filter = kwargs.get("document_ids_filter")
@@ -107,7 +121,7 @@ class Jieba(BaseKeyword):
if document_ids_filter:
segment_query_stmt = segment_query_stmt.where(DocumentSegment.document_id.in_(document_ids_filter))
segments = db.session.scalars(segment_query_stmt).all()
segments = session.scalars(segment_query_stmt).all()
segment_map = {segment.index_node_id: segment for segment in segments}
for chunk_index in sorted_chunk_indices:
segment = segment_map.get(chunk_index)
@@ -128,39 +142,43 @@ class Jieba(BaseKeyword):
return documents
@override
def delete(self):
def delete(self, *, session: Session):
lock_name = f"keyword_indexing_lock_{self.dataset.id}"
with redis_client.lock(lock_name, timeout=600):
dataset_keyword_table = self.dataset.dataset_keyword_table
dataset_keyword_table = self.dataset.get_dataset_keyword_table(session=session)
if dataset_keyword_table:
db.session.delete(dataset_keyword_table)
db.session.commit()
session.delete(dataset_keyword_table)
session.commit()
if dataset_keyword_table.data_source_type != "database":
file_key = "keyword_files/" + self.dataset.tenant_id + "/" + self.dataset.id + ".txt"
storage.delete(file_key)
def _save_dataset_keyword_table(self, keyword_table: dict[str, set[str]] | None):
def _save_dataset_keyword_table(self, keyword_table: dict[str, set[str]] | None, session: Session):
keyword_table_dict = {
"__type__": "keyword_table",
"__data__": {"index_id": self.dataset.id, "summary": None, "table": keyword_table},
}
dataset_keyword_table = self.dataset.dataset_keyword_table
dataset_keyword_table = session.scalar(
select(DatasetKeywordTable).where(DatasetKeywordTable.dataset_id == self.dataset.id)
)
keyword_data_source_type = dataset_keyword_table.data_source_type if dataset_keyword_table else "file"
if keyword_data_source_type == "database":
if dataset_keyword_table is None:
return
dataset_keyword_table.keyword_table = dumps_with_sets(keyword_table_dict)
db.session.commit()
session.flush()
else:
file_key = "keyword_files/" + self.dataset.tenant_id + "/" + self.dataset.id + ".txt"
if storage.exists(file_key):
storage.delete(file_key)
storage.save(file_key, dumps_with_sets(keyword_table_dict).encode("utf-8"))
def _get_dataset_keyword_table(self) -> dict[str, set[str]] | None:
dataset_keyword_table = self.dataset.dataset_keyword_table
def _get_dataset_keyword_table(self, session: Session) -> dict[str, set[str]] | None:
dataset_keyword_table = session.scalar(
select(DatasetKeywordTable).where(DatasetKeywordTable.dataset_id == self.dataset.id)
)
if dataset_keyword_table:
keyword_table_dict = dataset_keyword_table.keyword_table_dict
keyword_table_dict = dataset_keyword_table.get_keyword_table_dict(session=session)
if keyword_table_dict:
data: Any = keyword_table_dict["__data__"]
return dict(data["table"])
@@ -178,8 +196,8 @@ class Jieba(BaseKeyword):
"__data__": {"index_id": self.dataset.id, "summary": None, "table": {}},
}
)
db.session.add(dataset_keyword_table)
db.session.commit()
session.add(dataset_keyword_table)
session.flush()
return {}
@@ -228,25 +246,25 @@ class Jieba(BaseKeyword):
return sorted_chunk_indices[:k]
def _update_segment_keywords(self, dataset_id: str, node_id: str, keywords: list[str]):
def _update_segment_keywords(self, dataset_id: str, node_id: str, keywords: list[str], session: Session):
stmt = select(DocumentSegment).where(
DocumentSegment.dataset_id == dataset_id, DocumentSegment.index_node_id == node_id
)
document_segment = db.session.scalar(stmt)
document_segment = session.scalar(stmt)
if document_segment:
document_segment.keywords = keywords
db.session.add(document_segment)
db.session.commit()
session.add(document_segment)
session.flush()
def create_segment_keywords(self, node_id: str, keywords: list[str]):
keyword_table = self._get_dataset_keyword_table()
self._update_segment_keywords(self.dataset.id, node_id, keywords)
def create_segment_keywords(self, node_id: str, keywords: list[str], session: Session):
keyword_table = self._get_dataset_keyword_table(session)
self._update_segment_keywords(self.dataset.id, node_id, keywords, session)
keyword_table = self._add_text_to_keyword_table(keyword_table or {}, node_id, keywords)
self._save_dataset_keyword_table(keyword_table)
self._save_dataset_keyword_table(keyword_table, session)
def multi_create_segment_keywords(self, pre_segment_data_list: list[PreSegmentData]):
def multi_create_segment_keywords(self, pre_segment_data_list: list[PreSegmentData], session: Session):
keyword_table_handler = JiebaKeywordTableHandler()
keyword_table = self._get_dataset_keyword_table()
keyword_table = self._get_dataset_keyword_table(session)
for pre_segment_data in pre_segment_data_list:
segment = pre_segment_data["segment"]
if pre_segment_data["keywords"]:
@@ -264,12 +282,12 @@ class Jieba(BaseKeyword):
keyword_table = self._add_text_to_keyword_table(
keyword_table or {}, segment.index_node_id, list(keywords)
)
self._save_dataset_keyword_table(keyword_table)
self._save_dataset_keyword_table(keyword_table, session)
def update_segment_keywords_index(self, node_id: str, keywords: list[str]):
keyword_table = self._get_dataset_keyword_table()
def update_segment_keywords_index(self, node_id: str, keywords: list[str], session: Session):
keyword_table = self._get_dataset_keyword_table(session)
keyword_table = self._add_text_to_keyword_table(keyword_table or {}, node_id, keywords)
self._save_dataset_keyword_table(keyword_table)
self._save_dataset_keyword_table(keyword_table, session)
def set_orjson_default(obj: Any):
@@ -3,6 +3,8 @@ from __future__ import annotations
from abc import ABC, abstractmethod
from typing import Any
from sqlalchemy.orm import Session
from core.rag.models.document import Document
from models.dataset import Dataset
@@ -12,35 +14,35 @@ class BaseKeyword(ABC):
self.dataset = dataset
@abstractmethod
def create(self, texts: list[Document], **kwargs) -> BaseKeyword:
def create(self, texts: list[Document], session: Session, **kwargs: Any) -> BaseKeyword:
raise NotImplementedError
@abstractmethod
def add_texts(self, texts: list[Document], **kwargs):
def add_texts(self, texts: list[Document], session: Session, **kwargs: Any):
raise NotImplementedError
@abstractmethod
def text_exists(self, id: str) -> bool:
def text_exists(self, id: str, *, session: Session) -> bool:
raise NotImplementedError
@abstractmethod
def delete_by_ids(self, ids: list[str]):
def delete_by_ids(self, ids: list[str], session: Session, **kwargs: Any):
raise NotImplementedError
@abstractmethod
def delete(self):
def delete(self, *, session: Session):
raise NotImplementedError
@abstractmethod
def search(self, query: str, **kwargs: Any) -> list[Document]:
def search(self, query: str, *, session: Session, **kwargs: Any) -> list[Document]:
raise NotImplementedError
def _filter_duplicate_texts(self, texts: list[Document]) -> list[Document]:
def _filter_duplicate_texts(self, texts: list[Document], *, session: Session) -> list[Document]:
for text in texts.copy():
if text.metadata is None:
continue
doc_id = text.metadata["doc_id"]
exists_duplicate_node = self.text_exists(doc_id)
exists_duplicate_node = self.text_exists(doc_id, session=session)
if exists_duplicate_node:
texts.remove(text)
@@ -1,5 +1,7 @@
from typing import Any
from sqlalchemy.orm import Session
from configs import dify_config
from core.rag.datasource.keyword.keyword_base import BaseKeyword
from core.rag.datasource.keyword.keyword_type import KeyWordType
@@ -27,23 +29,23 @@ class Keyword:
case _:
raise ValueError(f"Keyword store {keyword_type} is not supported.")
def create(self, texts: list[Document], **kwargs):
self._keyword_processor.create(texts, **kwargs)
def create(self, texts: list[Document], session: Session, **kwargs: Any):
self._keyword_processor.create(texts, session, **kwargs)
def add_texts(self, texts: list[Document], **kwargs):
self._keyword_processor.add_texts(texts, **kwargs)
def add_texts(self, texts: list[Document], session: Session, **kwargs: Any):
self._keyword_processor.add_texts(texts, session, **kwargs)
def text_exists(self, id: str) -> bool:
return self._keyword_processor.text_exists(id)
def text_exists(self, id: str, *, session: Session) -> bool:
return self._keyword_processor.text_exists(id, session=session)
def delete_by_ids(self, ids: list[str]):
self._keyword_processor.delete_by_ids(ids)
def delete_by_ids(self, ids: list[str], session: Session, **kwargs: Any):
self._keyword_processor.delete_by_ids(ids, session, **kwargs)
def delete(self):
self._keyword_processor.delete()
def delete(self, *, session: Session):
self._keyword_processor.delete(session=session)
def search(self, query: str, **kwargs: Any) -> list[Document]:
return self._keyword_processor.search(query, **kwargs)
def search(self, query: str, *, session: Session, **kwargs: Any) -> list[Document]:
return self._keyword_processor.search(query, session=session, **kwargs)
def __getattr__(self, name):
if self._keyword_processor is not None:
+172 -150
View File
@@ -12,7 +12,6 @@ from sqlalchemy.orm import Session, load_only
from configs import dify_config
from core.app.file_access import grant_upload_file_access
from core.db.session_factory import session_factory
from core.model_manager import ModelManager
from core.rag.data_post_processor.data_post_processor import DataPostProcessor, RerankingModelDict, WeightsDict
from core.rag.datasource.keyword.keyword_factory import Keyword
@@ -303,9 +302,13 @@ class RetrievalService:
keyword = Keyword(dataset=dataset)
documents = keyword.search(
cls.escape_query_for_search(query), top_k=top_k, document_ids_filter=document_ids_filter
)
with Session(db.engine) as session:
documents = keyword.search(
cls.escape_query_for_search(query),
session=session,
top_k=top_k,
document_ids_filter=document_ids_filter,
)
all_documents.extend(documents)
except Exception as e:
logger.error(e, exc_info=True)
@@ -333,7 +336,6 @@ class RetrievalService:
if not dataset:
raise ValueError("dataset not found")
vector = Vector(dataset=dataset)
documents = []
# Hybrid search merges keyword / full-text / vector hits and then reranks
# (weighted fusion or reranking model). Applying the user score threshold at
@@ -342,29 +344,31 @@ class RetrievalService:
embedding_score_threshold = (
0.0 if retrieval_method == RetrievalMethod.HYBRID_SEARCH else score_threshold
)
if query_type == QueryType.TEXT_QUERY:
documents.extend(
vector.search_by_vector(
query,
search_type="similarity_score_threshold",
top_k=top_k,
score_threshold=embedding_score_threshold,
filter={"group_id": [dataset.id]},
document_ids_filter=document_ids_filter,
with Session(db.engine) as session:
vector = Vector(dataset=dataset, session=session)
if query_type == QueryType.TEXT_QUERY:
documents.extend(
vector.search_by_vector(
query,
search_type="similarity_score_threshold",
top_k=top_k,
score_threshold=embedding_score_threshold,
filter={"group_id": [dataset.id]},
document_ids_filter=document_ids_filter,
)
)
)
if query_type == QueryType.IMAGE_QUERY:
if not dataset.is_multimodal:
return
documents.extend(
vector.search_by_file(
file_id=query,
top_k=top_k,
score_threshold=embedding_score_threshold,
filter={"group_id": [dataset.id]},
document_ids_filter=document_ids_filter,
if query_type == QueryType.IMAGE_QUERY:
if not dataset.is_multimodal:
return
documents.extend(
vector.search_by_file(
file_id=query,
top_k=top_k,
score_threshold=embedding_score_threshold,
filter={"group_id": [dataset.id]},
document_ids_filter=document_ids_filter,
)
)
)
if documents:
if (
@@ -373,18 +377,37 @@ class RetrievalService:
and reranking_model["reranking_provider_name"]
and retrieval_method == RetrievalMethod.SEMANTIC_SEARCH
):
data_post_processor = DataPostProcessor(
str(dataset.tenant_id), str(RerankMode.RERANKING_MODEL), reranking_model, None, False
)
if dataset.is_multimodal:
model_manager = ModelManager.for_tenant(tenant_id=dataset.tenant_id)
is_support_vision = model_manager.check_model_support_vision(
tenant_id=dataset.tenant_id,
provider=reranking_model["reranking_provider_name"],
model=reranking_model["reranking_model_name"],
model_type=ModelType.RERANK,
with Session(db.engine) as rerank_session:
data_post_processor = DataPostProcessor(
str(dataset.tenant_id),
str(RerankMode.RERANKING_MODEL),
reranking_model,
None,
False,
session=rerank_session,
)
if is_support_vision:
if dataset.is_multimodal:
model_manager = ModelManager.for_tenant(tenant_id=dataset.tenant_id)
is_support_vision = model_manager.check_model_support_vision(
tenant_id=dataset.tenant_id,
provider=reranking_model["reranking_provider_name"],
model=reranking_model["reranking_model_name"],
model_type=ModelType.RERANK,
)
if is_support_vision:
all_documents.extend(
data_post_processor.invoke(
query=query,
documents=documents,
score_threshold=score_threshold,
top_n=len(documents),
query_type=query_type,
)
)
else:
# not effective, return original documents
all_documents.extend(documents)
else:
all_documents.extend(
data_post_processor.invoke(
query=query,
@@ -394,19 +417,6 @@ class RetrievalService:
query_type=query_type,
)
)
else:
# not effective, return original documents
all_documents.extend(documents)
else:
all_documents.extend(
data_post_processor.invoke(
query=query,
documents=documents,
score_threshold=score_threshold,
top_n=len(documents),
query_type=query_type,
)
)
else:
all_documents.extend(documents)
except Exception as e:
@@ -434,7 +444,8 @@ class RetrievalService:
if not dataset:
raise ValueError("dataset not found")
vector_processor = Vector(dataset=dataset)
with Session(db.engine) as session:
vector_processor = Vector(dataset=dataset, session=session)
documents = vector_processor.search_by_full_text(
cls.escape_query_for_search(query), top_k=top_k, document_ids_filter=document_ids_filter
@@ -446,17 +457,23 @@ class RetrievalService:
and reranking_model["reranking_provider_name"]
and retrieval_method == RetrievalMethod.FULL_TEXT_SEARCH
):
data_post_processor = DataPostProcessor(
str(dataset.tenant_id), str(RerankMode.RERANKING_MODEL), reranking_model, None, False
)
all_documents.extend(
data_post_processor.invoke(
query=query,
documents=documents,
score_threshold=score_threshold,
top_n=len(documents),
with Session(db.engine) as rerank_session:
data_post_processor = DataPostProcessor(
str(dataset.tenant_id),
str(RerankMode.RERANKING_MODEL),
reranking_model,
None,
False,
session=rerank_session,
)
all_documents.extend(
data_post_processor.invoke(
query=query,
documents=documents,
score_threshold=score_threshold,
top_n=len(documents),
)
)
)
else:
all_documents.extend(documents)
except Exception as e:
@@ -468,7 +485,7 @@ class RetrievalService:
return query.replace('"', '\\"')
@classmethod
def format_retrieval_documents(cls, documents: list[Document]) -> list[RetrievalSegments]:
def format_retrieval_documents(cls, session: Session, documents: list[Document]) -> list[RetrievalSegments]:
"""Format retrieval documents with optimized batch processing"""
if not documents:
return []
@@ -482,7 +499,7 @@ class RetrievalService:
# Batch query dataset documents
dataset_documents = {
doc.id: doc
for doc in db.session.scalars(
for doc in session.scalars(
select(DatasetDocument)
.where(DatasetDocument.id.in_(document_ids))
.options(load_only(DatasetDocument.id, DatasetDocument.doc_form, DatasetDocument.dataset_id))
@@ -558,84 +575,83 @@ class RetrievalService:
doc_segment_map: dict[str, list[str]] = {}
segment_summary_map: dict[str, str] = {} # Map segment_id to summary content
with session_factory.create_session() as session:
attachments = cls.get_segment_attachment_infos(image_doc_ids, session)
attachments = cls.get_segment_attachment_infos(image_doc_ids, session)
for attachment in attachments:
segment_ids.append(attachment["segment_id"])
if attachment["segment_id"] in attachment_map:
attachment_map[attachment["segment_id"]].append(attachment["attachment_info"])
else:
attachment_map[attachment["segment_id"]] = [attachment["attachment_info"]]
if attachment["segment_id"] in doc_segment_map:
doc_segment_map[attachment["segment_id"]].append(attachment["attachment_id"])
else:
doc_segment_map[attachment["segment_id"]] = [attachment["attachment_id"]]
for attachment in attachments:
segment_ids.append(attachment["segment_id"])
if attachment["segment_id"] in attachment_map:
attachment_map[attachment["segment_id"]].append(attachment["attachment_info"])
else:
attachment_map[attachment["segment_id"]] = [attachment["attachment_info"]]
if attachment["segment_id"] in doc_segment_map:
doc_segment_map[attachment["segment_id"]].append(attachment["attachment_id"])
else:
doc_segment_map[attachment["segment_id"]] = [attachment["attachment_id"]]
child_chunk_stmt = select(ChildChunk).where(ChildChunk.index_node_id.in_(child_index_node_ids))
child_index_nodes = session.execute(child_chunk_stmt).scalars().all()
child_chunk_stmt = select(ChildChunk).where(ChildChunk.index_node_id.in_(child_index_node_ids))
child_index_nodes = session.execute(child_chunk_stmt).scalars().all()
for i in child_index_nodes:
assert i.index_node_id
segment_ids.append(i.segment_id)
if i.segment_id in child_chunk_map:
child_chunk_map[i.segment_id].append(i)
else:
child_chunk_map[i.segment_id] = [i]
if i.segment_id in doc_segment_map:
doc_segment_map[i.segment_id].append(i.index_node_id)
else:
doc_segment_map[i.segment_id] = [i.index_node_id]
for i in child_index_nodes:
assert i.index_node_id
segment_ids.append(i.segment_id)
if i.segment_id in child_chunk_map:
child_chunk_map[i.segment_id].append(i)
else:
child_chunk_map[i.segment_id] = [i]
if i.segment_id in doc_segment_map:
doc_segment_map[i.segment_id].append(i.index_node_id)
else:
doc_segment_map[i.segment_id] = [i.index_node_id]
if index_node_ids:
document_segment_stmt = select(DocumentSegment).where(
DocumentSegment.enabled == True,
DocumentSegment.status == "completed",
DocumentSegment.index_node_id.in_(index_node_ids),
if index_node_ids:
document_segment_stmt = select(DocumentSegment).where(
DocumentSegment.enabled == True,
DocumentSegment.status == "completed",
DocumentSegment.index_node_id.in_(index_node_ids),
)
index_node_segments = session.execute(document_segment_stmt).scalars().all()
for index_node_segment in index_node_segments:
assert index_node_segment.index_node_id
doc_segment_map[index_node_segment.id] = [index_node_segment.index_node_id]
if segment_ids:
document_segment_stmt = select(DocumentSegment).where(
DocumentSegment.enabled == True,
DocumentSegment.status == "completed",
DocumentSegment.id.in_(segment_ids),
)
segments = session.execute(document_segment_stmt).scalars().all() # type: ignore
if index_node_segments:
segments.extend(index_node_segments)
# Handle summary documents: query segments by original_chunk_id
if summary_segment_ids:
summary_segment_ids_list = list(summary_segment_ids)
summary_segment_stmt = select(DocumentSegment).where(
DocumentSegment.enabled == True,
DocumentSegment.status == "completed",
DocumentSegment.id.in_(summary_segment_ids_list),
)
summary_segments = session.execute(summary_segment_stmt).scalars().all() # type: ignore
segments.extend(summary_segments)
# Add summary segment IDs to segment_ids for summary query
for seg in summary_segments:
if seg.id not in segment_ids:
segment_ids.append(seg.id)
# Batch query summaries for segments retrieved via summary (only enabled summaries)
if summary_segment_ids:
summaries = session.scalars(
select(DocumentSegmentSummary).where(
DocumentSegmentSummary.chunk_id.in_(list(summary_segment_ids)),
DocumentSegmentSummary.status == "completed",
DocumentSegmentSummary.enabled.is_(True), # Only retrieve enabled summaries
)
index_node_segments = session.execute(document_segment_stmt).scalars().all()
for index_node_segment in index_node_segments:
assert index_node_segment.index_node_id
doc_segment_map[index_node_segment.id] = [index_node_segment.index_node_id]
if segment_ids:
document_segment_stmt = select(DocumentSegment).where(
DocumentSegment.enabled == True,
DocumentSegment.status == "completed",
DocumentSegment.id.in_(segment_ids),
)
segments = session.execute(document_segment_stmt).scalars().all() # type: ignore
if index_node_segments:
segments.extend(index_node_segments)
# Handle summary documents: query segments by original_chunk_id
if summary_segment_ids:
summary_segment_ids_list = list(summary_segment_ids)
summary_segment_stmt = select(DocumentSegment).where(
DocumentSegment.enabled == True,
DocumentSegment.status == "completed",
DocumentSegment.id.in_(summary_segment_ids_list),
)
summary_segments = session.execute(summary_segment_stmt).scalars().all() # type: ignore
segments.extend(summary_segments)
# Add summary segment IDs to segment_ids for summary query
for seg in summary_segments:
if seg.id not in segment_ids:
segment_ids.append(seg.id)
# Batch query summaries for segments retrieved via summary (only enabled summaries)
if summary_segment_ids:
summaries = session.scalars(
select(DocumentSegmentSummary).where(
DocumentSegmentSummary.chunk_id.in_(list(summary_segment_ids)),
DocumentSegmentSummary.status == "completed",
DocumentSegmentSummary.enabled.is_(True), # Only retrieve enabled summaries
)
).all()
for summary in summaries:
if summary.summary_content:
segment_summary_map[summary.chunk_id] = summary.summary_content
).all()
for summary in summaries:
if summary.summary_content:
segment_summary_map[summary.chunk_id] = summary.summary_content
include_segment_ids = set()
segment_child_map: dict[str, SegmentChildMapDetail] = {}
@@ -774,7 +790,7 @@ class RetrievalService:
return sorted(result, key=lambda x: x.score if x.score is not None else 0.0, reverse=True)
except Exception as e:
db.session.rollback()
session.rollback()
raise e
@trace_span()
@@ -882,9 +898,6 @@ class RetrievalService:
if attachment_id and reranking_mode == RerankMode.WEIGHTED_SCORE:
all_documents.extend(all_documents_item)
all_documents_item = self._deduplicate_documents(all_documents_item)
data_post_processor = DataPostProcessor(
str(dataset.tenant_id), reranking_mode, reranking_model, weights, False
)
if query:
rerank_query = query
@@ -894,17 +907,26 @@ class RetrievalService:
query_type = QueryType.IMAGE_QUERY
else:
return
all_documents_item = data_post_processor.invoke(
query=rerank_query,
documents=all_documents_item,
score_threshold=score_threshold,
top_n=top_k,
query_type=query_type,
)
if not data_post_processor.rerank_runner and score_threshold:
all_documents_item = self._filter_documents_by_vector_score_threshold(
all_documents_item, score_threshold
with Session(db.engine) as rerank_session:
data_post_processor = DataPostProcessor(
str(dataset.tenant_id),
reranking_mode,
reranking_model,
weights,
False,
session=rerank_session,
)
all_documents_item = data_post_processor.invoke(
query=rerank_query,
documents=all_documents_item,
score_threshold=score_threshold,
top_n=top_k,
query_type=query_type,
)
if not data_post_processor.rerank_runner and score_threshold:
all_documents_item = self._filter_documents_by_vector_score_threshold(
all_documents_item, score_threshold
)
all_documents.extend(all_documents_item)
@@ -5,6 +5,7 @@ from abc import ABC, abstractmethod
from typing import Any, override
from sqlalchemy import select
from sqlalchemy.orm import Session
from configs import dify_config
from core.model_manager import ModelManager
@@ -15,7 +16,6 @@ from core.rag.embedding.cached_embedding import CacheEmbedding
from core.rag.embedding.embedding_base import Embeddings
from core.rag.index_processor.constant.doc_type import DocType
from core.rag.models.document import Document
from extensions.ext_database import db
from extensions.ext_redis import redis_client
from extensions.ext_storage import storage
from extensions.otel import trace_span
@@ -99,7 +99,7 @@ class _LazyEmbeddings(Embeddings):
class Vector:
def __init__(self, dataset: Dataset, attributes: list | None = None):
def __init__(self, dataset: Dataset, attributes: list | None = None, *, session: Session):
if attributes is None:
# `is_summary` and `original_chunk_id` are stored on summary vectors
# by `SummaryIndexService` and read back by `RetrievalService` to
@@ -120,14 +120,15 @@ class Vector:
]
self._dataset = dataset
# Use a lazy proxy so cleanup paths (delete_by_ids / delete / text_exists)
# never transitively trigger billing API calls during ``Vector(dataset)``
# never transitively trigger billing API calls during ``Vector(dataset, session=...)``
# construction. The real embedding model is materialized only when an
# ``embed_*`` method is actually invoked (i.e. create / search paths).
self._embeddings: Embeddings = _LazyEmbeddings(dataset)
self._attributes = attributes
self._vector_processor = self._init_vector()
self._session = session
self._vector_processor = self._init_vector(session=session)
def _init_vector(self) -> BaseVector:
def _init_vector(self, *, session: Session) -> BaseVector:
vector_type = dify_config.VECTOR_STORE
if self._dataset.index_struct_dict:
@@ -137,7 +138,7 @@ class Vector:
stmt = select(Whitelist).where(
Whitelist.tenant_id == self._dataset.tenant_id, Whitelist.category == "vector_db"
)
whitelist = db.session.scalars(stmt).one_or_none()
whitelist = session.scalars(stmt).one_or_none()
if whitelist:
vector_type = VectorType.TIDB_ON_QDRANT
@@ -194,7 +195,7 @@ class Vector:
# Batch query all upload files to avoid N+1 queries
attachment_ids = [doc.metadata["doc_id"] for doc in batch]
stmt = select(UploadFile).where(UploadFile.id.in_(attachment_ids))
upload_files = db.session.scalars(stmt).all()
upload_files = self._session.scalars(stmt).all()
upload_file_map = {str(f.id): f for f in upload_files}
file_base64_list = []
@@ -252,7 +253,7 @@ class Vector:
return self._vector_processor.search_by_vector(query_vector, **kwargs)
def search_by_file(self, file_id: str, **kwargs: Any) -> list[Document]:
upload_file: UploadFile | None = db.session.get(UploadFile, file_id)
upload_file: UploadFile | None = self._session.get(UploadFile, file_id)
if not upload_file:
return []
+41 -30
View File
@@ -4,11 +4,11 @@ from collections.abc import Sequence
from typing import Any
from sqlalchemy import delete, func, select
from sqlalchemy.orm import Session
from core.model_manager import ModelManager
from core.rag.index_processor.constant.index_type import IndexTechniqueType
from core.rag.models.document import AttachmentDocument, Document
from extensions.ext_database import db
from graphon.model_runtime.entities.model_entities import ModelType
from models.dataset import ChildChunk, Dataset, DocumentSegment, SegmentAttachmentBinding
from models.enums import SegmentType
@@ -45,8 +45,11 @@ class DatasetDocumentStore:
@property
def docs(self) -> dict[str, Document]:
raise ValueError("session is required; use get_docs(session)")
def get_docs(self, session: Session) -> dict[str, Document]:
stmt = select(DocumentSegment).where(DocumentSegment.dataset_id == self._dataset.id)
document_segments = db.session.scalars(stmt).all()
document_segments = session.scalars(stmt).all()
output = {}
for document_segment in document_segments:
@@ -64,8 +67,14 @@ class DatasetDocumentStore:
return output
def add_documents(self, docs: Sequence[Document], allow_update: bool = True, save_child: bool = False):
max_position = db.session.scalar(
def add_documents(
self,
docs: Sequence[Document],
session: Session,
allow_update: bool = True,
save_child: bool = False,
):
max_position = session.scalar(
select(func.max(DocumentSegment.position)).where(DocumentSegment.document_id == self._document_id)
)
@@ -94,7 +103,7 @@ class DatasetDocumentStore:
if doc.metadata is None:
raise ValueError("doc.metadata must be a dict")
segment_document = self.get_document_segment(doc_id=doc.metadata["doc_id"])
segment_document = self.get_document_segment(doc_id=doc.metadata["doc_id"], session=session)
# NOTE: doc could already exist in the store, but we overwrite it
if not allow_update and segment_document:
@@ -121,10 +130,10 @@ class DatasetDocumentStore:
if doc.metadata.get("answer"):
segment_document.answer = doc.metadata.pop("answer", "")
db.session.add(segment_document)
db.session.flush()
session.add(segment_document)
session.flush()
self.add_multimodel_documents_binding(
segment_id=segment_document.id, multimodel_documents=doc.attachments
segment_id=segment_document.id, multimodel_documents=doc.attachments, session=session
)
if save_child:
if doc.children:
@@ -143,7 +152,7 @@ class DatasetDocumentStore:
type=SegmentType.AUTOMATIC,
created_by=self._user_id,
)
db.session.add(child_segment)
session.add(child_segment)
else:
segment_document.content = doc.page_content
if doc.metadata.get("answer"):
@@ -152,11 +161,11 @@ class DatasetDocumentStore:
segment_document.word_count = len(doc.page_content)
segment_document.tokens = tokens
self.add_multimodel_documents_binding(
segment_id=segment_document.id, multimodel_documents=doc.attachments
segment_id=segment_document.id, multimodel_documents=doc.attachments, session=session
)
if save_child and doc.children:
# delete the existing child chunks
db.session.execute(
session.execute(
delete(ChildChunk).where(
ChildChunk.tenant_id == self._dataset.tenant_id,
ChildChunk.dataset_id == self._dataset.id,
@@ -180,17 +189,17 @@ class DatasetDocumentStore:
type=SegmentType.AUTOMATIC,
created_by=self._user_id,
)
db.session.add(child_segment)
session.add(child_segment)
db.session.commit()
session.flush()
def document_exists(self, doc_id: str) -> bool:
def document_exists(self, doc_id: str, session: Session) -> bool:
"""Check if document exists."""
result = self.get_document_segment(doc_id)
result = self.get_document_segment(doc_id, session=session)
return result is not None
def get_document(self, doc_id: str, raise_error: bool = True) -> Document | None:
document_segment = self.get_document_segment(doc_id)
def get_document(self, doc_id: str, session: Session, raise_error: bool = True) -> Document | None:
document_segment = self.get_document_segment(doc_id, session=session)
if document_segment is None:
if raise_error:
@@ -208,8 +217,8 @@ class DatasetDocumentStore:
},
)
def delete_document(self, doc_id: str, raise_error: bool = True):
document_segment = self.get_document_segment(doc_id)
def delete_document(self, doc_id: str, session: Session, raise_error: bool = True):
document_segment = self.get_document_segment(doc_id, session=session)
if document_segment is None:
if raise_error:
@@ -217,37 +226,39 @@ class DatasetDocumentStore:
else:
return None
db.session.delete(document_segment)
db.session.commit()
session.delete(document_segment)
session.flush()
def set_document_hash(self, doc_id: str, doc_hash: str):
def set_document_hash(self, doc_id: str, doc_hash: str, session: Session):
"""Set the hash for a given doc_id."""
document_segment = self.get_document_segment(doc_id)
document_segment = self.get_document_segment(doc_id, session=session)
if document_segment is None:
return None
document_segment.index_node_hash = doc_hash
db.session.commit()
session.flush()
def get_document_hash(self, doc_id: str) -> str | None:
def get_document_hash(self, doc_id: str, session: Session) -> str | None:
"""Get the stored hash for a document, if it exists."""
document_segment = self.get_document_segment(doc_id)
document_segment = self.get_document_segment(doc_id, session=session)
if document_segment is None:
return None
data: str | None = document_segment.index_node_hash
return data
def get_document_segment(self, doc_id: str) -> DocumentSegment | None:
def get_document_segment(self, doc_id: str, session: Session) -> DocumentSegment | None:
stmt = select(DocumentSegment).where(
DocumentSegment.dataset_id == self._dataset.id, DocumentSegment.index_node_id == doc_id
)
document_segment = db.session.scalar(stmt)
document_segment = session.scalar(stmt)
return document_segment
def add_multimodel_documents_binding(self, segment_id: str, multimodel_documents: list[AttachmentDocument] | None):
def add_multimodel_documents_binding(
self, segment_id: str, multimodel_documents: list[AttachmentDocument] | None, session: Session
):
if multimodel_documents and self._document_id is not None:
for multimodel_document in multimodel_documents:
binding = SegmentAttachmentBinding(
@@ -257,4 +268,4 @@ class DatasetDocumentStore:
segment_id=segment_id,
attachment_id=multimodel_document.metadata["doc_id"],
)
db.session.add(binding)
session.add(binding)
+20 -5
View File
@@ -4,6 +4,8 @@ from pathlib import Path
from typing import Literal, overload
from urllib.parse import unquote
from sqlalchemy.orm import Session
from configs import dify_config
from core.file import remote_fetcher
from core.rag.extractor.csv_extractor import CSVExtractor
@@ -111,7 +113,12 @@ class ExtractProcessor:
@classmethod
def extract(
cls, extract_setting: ExtractSetting, is_automatic: bool = False, file_path: str | None = None
cls,
extract_setting: ExtractSetting,
is_automatic: bool = False,
file_path: str | None = None,
*,
session: Session | None = None,
) -> list[Document]:
if extract_setting.datasource_type == DatasourceType.FILE:
upload_file = extract_setting.upload_file
@@ -141,7 +148,9 @@ class ExtractProcessor:
)
elif file_extension == ".pdf":
assert upload_file is not None
extractor = PdfExtractor(file_path, upload_file.tenant_id, upload_file.created_by)
extractor = PdfExtractor(
file_path, upload_file.tenant_id, upload_file.created_by, session=session
)
elif file_extension in {".md", ".markdown", ".mdx"}:
extractor = (
UnstructuredMarkdownExtractor(file_path, unstructured_api_url, unstructured_api_key)
@@ -152,7 +161,9 @@ class ExtractProcessor:
extractor = HtmlExtractor(file_path)
elif file_extension == ".docx":
assert upload_file is not None
extractor = WordExtractor(file_path, upload_file.tenant_id, upload_file.created_by)
extractor = WordExtractor(
file_path, upload_file.tenant_id, upload_file.created_by, session=session
)
elif file_extension == ".doc":
extractor = UnstructuredWordExtractor(file_path, unstructured_api_url, unstructured_api_key)
elif file_extension == ".csv":
@@ -184,14 +195,18 @@ class ExtractProcessor:
)
elif file_extension == ".pdf":
assert upload_file is not None
extractor = PdfExtractor(file_path, upload_file.tenant_id, upload_file.created_by)
extractor = PdfExtractor(
file_path, upload_file.tenant_id, upload_file.created_by, session=session
)
elif file_extension in {".md", ".markdown", ".mdx"}:
extractor = MarkdownExtractor(file_path, autodetect_encoding=True)
elif file_extension in {".htm", ".html"}:
extractor = HtmlExtractor(file_path)
elif file_extension == ".docx":
assert upload_file is not None
extractor = WordExtractor(file_path, upload_file.tenant_id, upload_file.created_by)
extractor = WordExtractor(
file_path, upload_file.tenant_id, upload_file.created_by, session=session
)
elif file_extension == ".csv":
extractor = CSVExtractor(file_path, autodetect_encoding=True)
elif file_extension == ".epub":
+17 -3
View File
@@ -9,6 +9,7 @@ from typing import override
import pypdfium2
import pypdfium2.raw as pdfium_c
from sqlalchemy.orm import Session
from configs import dify_config
from core.rag.extractor.blob.blob import Blob
@@ -33,6 +34,7 @@ class PdfExtractor(BaseExtractor):
tenant_id: Workspace ID.
user_id: ID of the user performing the extraction.
file_cache_key: Optional cache key for the extracted text.
session: Session used to persist extracted images.
"""
# Magic bytes for image format detection: (magic_bytes, extension, mime_type)
@@ -48,13 +50,23 @@ class PdfExtractor(BaseExtractor):
(b"MM\x00+", "tiff", "image/tiff"),
)
MAX_MAGIC_LEN = max(len(m) for m, _, _ in IMAGE_FORMATS)
_session: Session | None
def __init__(self, file_path: str, tenant_id: str, user_id: str, file_cache_key: str | None = None):
def __init__(
self,
file_path: str,
tenant_id: str,
user_id: str,
file_cache_key: str | None = None,
*,
session: Session | None = None,
):
"""Initialize PdfExtractor."""
self._file_path = file_path
self._tenant_id = tenant_id
self._user_id = user_id
self._file_cache_key = file_cache_key
self._session = session
@override
def extract(self) -> list[Document]:
@@ -174,6 +186,8 @@ class PdfExtractor(BaseExtractor):
except Exception as e:
logger.warning("Failed to get objects from PDF page: %s", e)
if upload_files:
db.session.add_all(upload_files)
db.session.commit()
session = self._session or db.session
session.add_all(upload_files)
if self._session is None:
session.commit()
return "\n".join(image_content)
+13 -4
View File
@@ -18,6 +18,7 @@ from docx.oxml.ns import qn
from docx.table import Table
from docx.text.paragraph import Paragraph
from docx.text.run import Run
from sqlalchemy.orm import Session
from configs import dify_config
from core.file import remote_fetcher
@@ -38,16 +39,19 @@ class WordExtractor(BaseExtractor):
Args:
file_path: Path to the file to load.
session: Session used to persist extracted images.
"""
_closed: bool
_session: Session | None
def __init__(self, file_path: str, tenant_id: str, user_id: str):
def __init__(self, file_path: str, tenant_id: str, user_id: str, *, session: Session | None = None):
"""Initialize with file path."""
self._closed = False
self.file_path = file_path
self.tenant_id = tenant_id
self.user_id = user_id
self._session = session
if "~" in self.file_path:
self.file_path = os.path.expanduser(self.file_path)
@@ -112,8 +116,10 @@ class WordExtractor(BaseExtractor):
return bool(parsed.netloc) and bool(parsed.scheme)
def _extract_images_from_docx(self, doc):
session = self._session or db.session
image_count = 0
image_map = {}
upload_files: list[UploadFile] = []
base_url = dify_config.FILES_URL
for r_id, rel in doc.part.rels.items():
@@ -152,7 +158,7 @@ class WordExtractor(BaseExtractor):
used_by=self.user_id,
used_at=naive_utc_now(),
)
db.session.add(upload_file)
upload_files.append(upload_file)
image_map[r_id] = f"![image]({base_url}/files/{upload_file.id}/file-preview)"
else:
image_ext = rel.target_ref.split(".")[-1]
@@ -180,9 +186,12 @@ class WordExtractor(BaseExtractor):
used_by=self.user_id,
used_at=naive_utc_now(),
)
db.session.add(upload_file)
upload_files.append(upload_file)
image_map[rel.target_part] = f"![image]({base_url}/files/{upload_file.id}/file-preview)"
db.session.commit()
if upload_files:
session.add_all(upload_files)
if self._session is None:
session.commit()
return image_map
def _table_to_markdown(self, table, image_map):
+75 -80
View File
@@ -7,6 +7,7 @@ from typing import Any
from flask import current_app
from sqlalchemy import delete, func, select, update
from sqlalchemy.orm import Session
from core.db.session_factory import session_factory
from core.rag.index_processor.constant.index_type import IndexTechniqueType
@@ -61,75 +62,77 @@ class IndexProcessor:
chunks: Mapping[str, Any],
batch: Any,
summary_index_setting: SummaryIndexSettingDict | None = None,
*,
session: Session,
) -> IndexingResultDict:
with session_factory.create_session() as session:
document = session.scalar(select(Document).where(Document.id == document_id).limit(1))
if not document:
raise KnowledgeIndexNodeError(f"Document {document_id} not found.")
document = session.scalar(select(Document).where(Document.id == document_id).limit(1))
if not document:
raise KnowledgeIndexNodeError(f"Document {document_id} not found.")
dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1))
if not dataset:
raise KnowledgeIndexNodeError(f"Dataset {dataset_id} not found.")
dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1))
if not dataset:
raise KnowledgeIndexNodeError(f"Dataset {dataset_id} not found.")
dataset_name_value = dataset.name
document_name_value = document.name
created_at_value = document.created_at
if summary_index_setting is None:
summary_index_setting = dataset.summary_index_setting
index_node_ids = []
dataset_name_value = dataset.name
document_name_value = document.name
created_at_value = document.created_at
if summary_index_setting is None:
summary_index_setting = dataset.summary_index_setting
index_node_ids = []
index_processor = IndexProcessorFactory(dataset.chunk_structure).init_index_processor()
if original_document_id:
segments = session.scalars(
select(DocumentSegment).where(DocumentSegment.document_id == original_document_id)
).all()
if segments:
index_node_ids = [segment.index_node_id for segment in segments if segment.index_node_id]
index_processor = IndexProcessorFactory(dataset.chunk_structure).init_index_processor()
if original_document_id:
segments = session.scalars(
select(DocumentSegment).where(DocumentSegment.document_id == original_document_id)
).all()
if segments:
index_node_ids = [segment.index_node_id for segment in segments if segment.index_node_id]
indexing_start_at = time.perf_counter()
# The metadata reads above must not keep a transaction open across vector I/O.
session.commit()
# delete from vector index
if index_node_ids:
index_processor.clean(dataset, index_node_ids, with_keywords=True, delete_child_chunks=True)
index_processor.clean(
dataset, index_node_ids, with_keywords=True, delete_child_chunks=True, session=session
)
session.commit()
segment_delete_stmt = delete(DocumentSegment).where(DocumentSegment.document_id == original_document_id)
session.execute(segment_delete_stmt)
session.commit()
with session_factory.create_session() as session, session.begin():
if index_node_ids:
segment_delete_stmt = delete(DocumentSegment).where(DocumentSegment.document_id == original_document_id)
session.execute(segment_delete_stmt)
index_processor.index(dataset, document, chunks)
index_processor.index(dataset, document, chunks, session)
session.commit()
indexing_end_at = time.perf_counter()
with session_factory.create_session() as session, session.begin():
document.indexing_latency = indexing_end_at - indexing_start_at
document.indexing_status = "completed"
document.completed_at = datetime.datetime.now(datetime.UTC).replace(tzinfo=None)
document.word_count = (
session.scalar(
select(func.sum(DocumentSegment.word_count)).where(
DocumentSegment.document_id == document_id,
DocumentSegment.dataset_id == dataset_id,
)
)
) or 0
# Update need_summary based on dataset's summary_index_setting
if summary_index_setting and summary_index_setting.get("enable") is True:
document.need_summary = True
else:
document.need_summary = False
session.add(document)
# update document segment status
session.execute(
update(DocumentSegment)
.where(
document.indexing_latency = indexing_end_at - indexing_start_at
document.indexing_status = "completed"
document.completed_at = datetime.datetime.now(datetime.UTC).replace(tzinfo=None)
document.word_count = (
session.scalar(
select(func.sum(DocumentSegment.word_count)).where(
DocumentSegment.document_id == document_id,
DocumentSegment.dataset_id == dataset_id,
)
.values(
status="completed",
enabled=True,
completed_at=datetime.datetime.now(datetime.UTC).replace(tzinfo=None),
)
)
) or 0
# Update need_summary based on dataset's summary_index_setting
document.need_summary = bool(summary_index_setting and summary_index_setting.get("enable") is True)
session.add(document)
# update document segment status
session.execute(
update(DocumentSegment)
.where(
DocumentSegment.document_id == document_id,
DocumentSegment.dataset_id == dataset_id,
)
.values(
status="completed",
enabled=True,
completed_at=datetime.datetime.now(datetime.UTC).replace(tzinfo=None),
)
)
session.flush()
result: IndexingResultDict = {
"dataset_id": dataset_id,
@@ -149,25 +152,27 @@ class IndexProcessor:
document_id: str,
chunk_structure: str,
summary_index_setting: SummaryIndexSettingDict | None,
*,
session: Session,
) -> Preview:
doc_language = None
with session_factory.create_session() as session:
if document_id:
document = session.scalar(select(Document).where(Document.id == document_id).limit(1))
else:
document = None
if document_id:
document = session.scalar(select(Document).where(Document.id == document_id).limit(1))
else:
document = None
dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1))
if not dataset:
raise KnowledgeIndexNodeError(f"Dataset {dataset_id} not found.")
dataset = session.scalar(select(Dataset).where(Dataset.id == dataset_id).limit(1))
if not dataset:
raise KnowledgeIndexNodeError(f"Dataset {dataset_id} not found.")
if summary_index_setting is None:
summary_index_setting = dataset.summary_index_setting
if summary_index_setting is None:
summary_index_setting = dataset.summary_index_setting
if document:
doc_language = document.doc_language
indexing_technique = dataset.indexing_technique
tenant_id = dataset.tenant_id
if document:
doc_language = document.doc_language
indexing_technique = dataset.indexing_technique
tenant_id = dataset.tenant_id
session.commit()
preview_output = self.format_preview(chunk_structure, chunks)
if indexing_technique != IndexTechniqueType.HIGH_QUALITY:
@@ -194,23 +199,13 @@ class IndexProcessor:
"""Generate summary for a single chunk."""
if flask_app:
with flask_app.app_context():
if preview_item.content is not None:
# Set Flask application context in worker thread
summary, _ = ParagraphIndexProcessor.generate_summary(
tenant_id=tenant_id,
text=preview_item.content,
summary_index_setting=summary_index_setting,
document_language=doc_language,
)
if summary:
preview_item.summary = summary
else:
with session_factory.create_session() as worker_session:
summary, _ = ParagraphIndexProcessor.generate_summary(
tenant_id=tenant_id,
text=preview_item.content if preview_item.content is not None else "",
summary_index_setting=summary_index_setting,
document_language=doc_language,
session=worker_session,
)
if summary:
preview_item.summary = summary
@@ -12,6 +12,7 @@ from urllib.parse import unquote, urlparse
import httpx
from sqlalchemy import select
from sqlalchemy.orm import Session
from configs import dify_config
from core.entities.knowledge_entities import PreviewDetail
@@ -46,11 +47,13 @@ class BaseIndexProcessor(ABC):
"""Interface for extract files."""
@abstractmethod
def extract(self, extract_setting: ExtractSetting, **kwargs) -> list[Document]:
def extract(self, extract_setting: ExtractSetting, *, session: Session, **kwargs) -> list[Document]:
raise NotImplementedError
@abstractmethod
def transform(self, documents: list[Document], current_user: Account | None = None, **kwargs) -> list[Document]:
def transform(
self, documents: list[Document], current_user: Account | None = None, *, session: Session, **kwargs
) -> list[Document]:
raise NotImplementedError
@abstractmethod
@@ -60,6 +63,8 @@ class BaseIndexProcessor(ABC):
preview_texts: list[PreviewDetail],
summary_index_setting: SummaryIndexSettingDict,
doc_language: str | None = None,
*,
session: Session,
) -> list[PreviewDetail]:
"""
For each segment in preview_texts, generate a summary using LLM and attach it to the segment.
@@ -71,6 +76,7 @@ class BaseIndexProcessor(ABC):
preview_texts: List of preview details to generate summaries for
summary_index_setting: Summary index configuration
doc_language: Optional document language to ensure summary is generated in the correct language
session: SQLAlchemy session used for summary image lookups
"""
raise NotImplementedError
@@ -81,16 +87,20 @@ class BaseIndexProcessor(ABC):
documents: list[Document],
multimodal_documents: list[AttachmentDocument] | None = None,
with_keywords: bool = True,
*,
session: Session,
**kwargs,
) -> None:
raise NotImplementedError
@abstractmethod
def clean(self, dataset: Dataset, node_ids: list[str] | None, with_keywords: bool = True, **kwargs) -> None:
def clean(
self, dataset: Dataset, node_ids: list[str] | None, with_keywords: bool = True, *, session: Session, **kwargs
) -> None:
raise NotImplementedError
@abstractmethod
def index(self, dataset: Dataset, document: DatasetDocument, chunks: Any) -> None:
def index(self, dataset: Dataset, document: DatasetDocument, chunks: Any, session: Session) -> None:
raise NotImplementedError
@abstractmethod
@@ -136,7 +146,9 @@ class BaseIndexProcessor(ABC):
return character_splitter
def _get_content_files(self, document: Document, current_user: Account | None = None) -> list[AttachmentDocument]:
def _get_content_files(
self, document: Document, current_user: Account | None = None, *, session: Session
) -> list[AttachmentDocument]:
"""
Get the content files from the document.
"""
@@ -173,7 +185,7 @@ class BaseIndexProcessor(ABC):
if match:
if current_user:
tool_file_id = match.group(1)
upload_file_id = self._download_tool_file(tool_file_id, current_user)
upload_file_id = self._download_tool_file(tool_file_id, current_user, session=session)
if upload_file_id:
upload_file_id_list.append(upload_file_id)
continue
@@ -187,7 +199,7 @@ class BaseIndexProcessor(ABC):
# Get unique IDs for database query
unique_upload_file_ids = list(set(upload_file_id_list))
upload_files = db.session.scalars(select(UploadFile).where(UploadFile.id.in_(unique_upload_file_ids))).all()
upload_files = session.scalars(select(UploadFile).where(UploadFile.id.in_(unique_upload_file_ids))).all()
# Create a mapping from ID to UploadFile for quick lookup
upload_file_map = {upload_file.id: upload_file for upload_file in upload_files}
@@ -293,13 +305,13 @@ class BaseIndexProcessor(ABC):
logging.warning("Unexpected error downloading image from %s", image_url, exc_info=True)
return None
def _download_tool_file(self, tool_file_id: str, current_user: Account) -> str | None:
def _download_tool_file(self, tool_file_id: str, current_user: Account, *, session: Session) -> str | None:
"""
Download the tool file from the ID.
"""
from services.file_service import FileService
tool_file = db.session.get(ToolFile, tool_file_id)
tool_file = session.get(ToolFile, tool_file_id)
if not tool_file:
return None
blob = storage.load_once(tool_file.file_key)
@@ -5,14 +5,12 @@ import re
import uuid
from typing import Any, TypedDict, cast, override
from sqlalchemy.orm import Session
logger = logging.getLogger(__name__)
from sqlalchemy import select
from sqlalchemy.orm import Session
from core.app.file_access import DatabaseFileAccessController
from core.app.llm import deduct_llm_quota
from core.db.session_factory import session_factory
from core.entities.knowledge_entities import PreviewDetail
from core.llm_generator.prompts import DEFAULT_GENERATOR_SUMMARY_PROMPT
from core.model_manager import ModelInstance
@@ -30,7 +28,6 @@ from core.rag.index_processor.index_processor_base import BaseIndexProcessor, Su
from core.rag.models.document import AttachmentDocument, Document, MultimodalGeneralStructureChunk
from core.tools.utils.text_processing_utils import remove_leading_symbols
from core.workflow.file_reference import build_file_reference
from extensions.ext_database import db
from factories.file_factory import build_from_mapping
from graphon.file import File, FileTransferMethod, FileType, file_manager
from graphon.model_runtime.entities.llm_entities import LLMResult, LLMUsage
@@ -50,6 +47,9 @@ from models.dataset import Document as DatasetDocument
from services.account_service import AccountService
from services.summary_index_service import SummaryIndexService
logger = logging.getLogger(__name__)
_file_access_controller = DatabaseFileAccessController()
@@ -61,18 +61,21 @@ class ParagraphFormatPreviewDict(TypedDict):
class ParagraphIndexProcessor(BaseIndexProcessor):
@override
def extract(self, extract_setting: ExtractSetting, **kwargs) -> list[Document]:
def extract(self, extract_setting: ExtractSetting, *, session: Session, **kwargs) -> list[Document]:
text_docs = ExtractProcessor.extract(
extract_setting=extract_setting,
is_automatic=(
kwargs.get("process_rule_mode") == "automatic" or kwargs.get("process_rule_mode") == "hierarchical"
),
session=session,
)
return text_docs
@override
def transform(self, documents: list[Document], current_user: Account | None = None, **kwargs) -> list[Document]:
def transform(
self, documents: list[Document], current_user: Account | None = None, *, session: Session, **kwargs
) -> list[Document]:
process_rule = kwargs.get("process_rule")
if not process_rule:
raise ValueError("No process rule found.")
@@ -109,7 +112,9 @@ class ParagraphIndexProcessor(BaseIndexProcessor):
document_node.metadata["doc_id"] = doc_id
document_node.metadata["doc_hash"] = hash
multimodal_documents = (
self._get_content_files(document_node, current_user) if document_node.metadata else None
self._get_content_files(document_node, current_user, session=session)
if document_node.metadata
else None
)
if multimodal_documents:
document_node.attachments = multimodal_documents
@@ -128,10 +133,12 @@ class ParagraphIndexProcessor(BaseIndexProcessor):
documents: list[Document],
multimodal_documents: list[AttachmentDocument] | None = None,
with_keywords: bool = True,
*,
session: Session,
**kwargs,
) -> None:
if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
vector = Vector(dataset)
vector = Vector(dataset, session=session)
vector.create(documents)
if multimodal_documents and dataset.is_multimodal:
vector.create_multimodal(multimodal_documents)
@@ -140,12 +147,14 @@ class ParagraphIndexProcessor(BaseIndexProcessor):
keywords_list = kwargs.get("keywords_list")
keyword = Keyword(dataset)
if keywords_list and len(keywords_list) > 0:
keyword.add_texts(documents, keywords_list=keywords_list)
keyword.add_texts(documents, session, keywords_list=keywords_list)
else:
keyword.add_texts(documents)
keyword.add_texts(documents, session)
@override
def clean(self, dataset: Dataset, node_ids: list[str] | None, with_keywords: bool = True, **kwargs) -> None:
def clean(
self, dataset: Dataset, node_ids: list[str] | None, with_keywords: bool = True, *, session: Session, **kwargs
) -> None:
# Note: Summary indexes are now disabled (not deleted) when segments are disabled.
# This method is called for actual deletion scenarios (e.g., when segment is deleted).
# For disable operations, disable_summaries_for_segments is called directly in the task.
@@ -154,7 +163,7 @@ class ParagraphIndexProcessor(BaseIndexProcessor):
if delete_summaries:
if node_ids:
# Find segments by index_node_id
segments = db.session.scalars(
segments = session.scalars(
select(DocumentSegment).where(
DocumentSegment.dataset_id == dataset.id,
DocumentSegment.index_node_id.in_(node_ids),
@@ -162,13 +171,13 @@ class ParagraphIndexProcessor(BaseIndexProcessor):
).all()
segment_ids = [segment.id for segment in segments]
if segment_ids:
SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=segment_ids)
SummaryIndexService.delete_summaries_for_segments(dataset, segment_ids, session=session)
else:
# Delete all summaries for the dataset
SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=None)
SummaryIndexService.delete_summaries_for_segments(dataset, None, session=session)
if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
vector = Vector(dataset)
vector = Vector(dataset, session=session)
if node_ids:
vector.delete_by_ids(node_ids)
else:
@@ -177,12 +186,12 @@ class ParagraphIndexProcessor(BaseIndexProcessor):
if with_keywords:
keyword = Keyword(dataset)
if node_ids:
keyword.delete_by_ids(node_ids)
keyword.delete_by_ids(node_ids, session)
else:
keyword.delete()
keyword.delete(session=session)
@override
def index(self, dataset: Dataset, document: DatasetDocument, chunks: Any) -> None:
def index(self, dataset: Dataset, document: DatasetDocument, chunks: Any, session: Session) -> None:
documents: list[Any] = []
all_multimodal_documents: list[Any] = []
if isinstance(chunks, list):
@@ -194,7 +203,7 @@ class ParagraphIndexProcessor(BaseIndexProcessor):
"doc_hash": helper.generate_text_hash(content),
}
doc = Document(page_content=content, metadata=metadata)
attachments = self._get_content_files(doc)
attachments = self._get_content_files(doc, session=session)
if attachments:
doc.attachments = attachments
all_multimodal_documents.extend(attachments)
@@ -226,10 +235,11 @@ class ParagraphIndexProcessor(BaseIndexProcessor):
all_multimodal_documents.append(file_document)
doc.attachments = attachments
else:
account = AccountService.load_user(document.created_by, db.session())
with session_factory.create_session() as account_session:
account = AccountService.load_user(document.created_by, account_session)
if not account:
raise ValueError("Invalid account")
doc.attachments = self._get_content_files(doc, current_user=account)
doc.attachments = self._get_content_files(doc, current_user=account, session=session)
if doc.attachments:
all_multimodal_documents.extend(doc.attachments)
documents.append(doc)
@@ -237,15 +247,16 @@ class ParagraphIndexProcessor(BaseIndexProcessor):
# save node to document segment
doc_store = DatasetDocumentStore(dataset=dataset, user_id=document.created_by, document_id=document.id)
# add document segments
doc_store.add_documents(docs=documents, save_child=False)
doc_store.add_documents(docs=documents, save_child=False, session=session)
session.commit()
if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
vector = Vector(dataset)
vector = Vector(dataset, session=session)
vector.create(documents)
if all_multimodal_documents and dataset.is_multimodal:
vector.create_multimodal(all_multimodal_documents)
elif dataset.indexing_technique == IndexTechniqueType.ECONOMY:
keyword = Keyword(dataset)
keyword.add_texts(documents)
keyword.add_texts(documents, session)
@override
def format_preview(self, chunks: Any) -> ParagraphFormatPreviewDict:
@@ -269,6 +280,8 @@ class ParagraphIndexProcessor(BaseIndexProcessor):
preview_texts: list[PreviewDetail],
summary_index_setting: SummaryIndexSettingDict,
doc_language: str | None = None,
*,
session: Session,
) -> list[PreviewDetail]:
"""
For each segment, concurrently call generate_summary to generate a summary
@@ -291,15 +304,25 @@ class ParagraphIndexProcessor(BaseIndexProcessor):
if flask_app:
# Ensure Flask app context in worker thread
with flask_app.app_context():
summary, _ = self.generate_summary(
tenant_id, preview.content, summary_index_setting, document_language=doc_language
)
with session_factory.create_session() as worker_session:
summary, _ = self.generate_summary(
tenant_id,
preview.content,
summary_index_setting,
document_language=doc_language,
session=worker_session,
)
preview.summary = summary
else:
# Fallback: try without app context (may fail)
summary, _ = self.generate_summary(
tenant_id, preview.content, summary_index_setting, document_language=doc_language
)
with session_factory.create_session() as worker_session:
summary, _ = self.generate_summary(
tenant_id,
preview.content,
summary_index_setting,
document_language=doc_language,
session=worker_session,
)
preview.summary = summary
# Generate summaries concurrently using ThreadPoolExecutor
@@ -354,6 +377,8 @@ class ParagraphIndexProcessor(BaseIndexProcessor):
summary_index_setting: SummaryIndexSettingDict | None = None,
segment_id: str | None = None,
document_language: str | None = None,
*,
session: Session,
) -> tuple[str, LLMUsage]:
"""
Generate summary for the given text using ModelInstance.invoke_llm and the default or custom summary prompt,
@@ -366,6 +391,7 @@ class ParagraphIndexProcessor(BaseIndexProcessor):
segment_id: Optional segment ID to fetch attachments from SegmentAttachmentBinding table
document_language: Optional document language (e.g., "Chinese", "English")
to ensure summary is generated in the correct language
session: SQLAlchemy session used for summary image lookups
Returns:
Tuple of (summary_content, llm_usage) where llm_usage is LLMUsage object
@@ -414,12 +440,12 @@ class ParagraphIndexProcessor(BaseIndexProcessor):
# First, try to get images from SegmentAttachmentBinding (preferred method)
if segment_id:
image_files = ParagraphIndexProcessor._extract_images_from_segment_attachments(
tenant_id, segment_id, db.session()
tenant_id, segment_id, session
)
# If no images from attachments, fall back to extracting from text
if not image_files:
image_files = ParagraphIndexProcessor._extract_images_from_text(tenant_id, text, db.session())
image_files = ParagraphIndexProcessor._extract_images_from_text(tenant_id, text, session)
# Build prompt messages
prompt_messages = []
@@ -6,6 +6,7 @@ import uuid
from typing import Any, TypedDict, override
from sqlalchemy import delete, select
from sqlalchemy.orm import Session
from configs import dify_config
from core.db.session_factory import session_factory
@@ -21,7 +22,6 @@ from core.rag.index_processor.constant.doc_type import DocType
from core.rag.index_processor.constant.index_type import IndexStructureType, IndexTechniqueType
from core.rag.index_processor.index_processor_base import BaseIndexProcessor, SummaryIndexSettingDict
from core.rag.models.document import AttachmentDocument, ChildDocument, Document, ParentChildStructureChunk
from extensions.ext_database import db
from libs import helper
from models import Account
from models.dataset import ChildChunk, Dataset, DatasetProcessRule, DocumentSegment
@@ -42,18 +42,21 @@ class ParentChildFormatPreviewDict(TypedDict):
class ParentChildIndexProcessor(BaseIndexProcessor):
@override
def extract(self, extract_setting: ExtractSetting, **kwargs) -> list[Document]:
def extract(self, extract_setting: ExtractSetting, *, session: Session, **kwargs) -> list[Document]:
text_docs = ExtractProcessor.extract(
extract_setting=extract_setting,
is_automatic=(
kwargs.get("process_rule_mode") == "automatic" or kwargs.get("process_rule_mode") == "hierarchical"
),
session=session,
)
return text_docs
@override
def transform(self, documents: list[Document], current_user: Account | None = None, **kwargs) -> list[Document]:
def transform(
self, documents: list[Document], current_user: Account | None = None, *, session: Session, **kwargs
) -> list[Document]:
process_rule = kwargs.get("process_rule")
if not process_rule:
raise ValueError("No process rule found.")
@@ -95,7 +98,7 @@ class ParentChildIndexProcessor(BaseIndexProcessor):
page_content = page_content
if len(page_content) > 0:
document_node.page_content = page_content
multimodel_documents = self._get_content_files(document_node, current_user)
multimodel_documents = self._get_content_files(document_node, current_user, session=session)
if multimodel_documents:
document_node.attachments = multimodel_documents
# parse document to child nodes
@@ -108,7 +111,7 @@ class ParentChildIndexProcessor(BaseIndexProcessor):
elif rules.parent_mode == ParentMode.FULL_DOC:
page_content = "\n".join([document.page_content for document in documents])
document = Document(page_content=page_content, metadata=documents[0].metadata)
multimodel_documents = self._get_content_files(document)
multimodel_documents = self._get_content_files(document, session=session)
if multimodel_documents:
document.attachments = multimodel_documents
# parse document to child nodes
@@ -135,10 +138,12 @@ class ParentChildIndexProcessor(BaseIndexProcessor):
documents: list[Document],
multimodal_documents: list[AttachmentDocument] | None = None,
with_keywords: bool = True,
*,
session: Session,
**kwargs,
) -> None:
if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
vector = Vector(dataset)
vector = Vector(dataset, session=session)
for document in documents:
child_documents = document.children
if child_documents:
@@ -150,7 +155,9 @@ class ParentChildIndexProcessor(BaseIndexProcessor):
vector.create_multimodal(multimodal_documents)
@override
def clean(self, dataset: Dataset, node_ids: list[str] | None, with_keywords: bool = True, **kwargs) -> None:
def clean(
self, dataset: Dataset, node_ids: list[str] | None, with_keywords: bool = True, *, session: Session, **kwargs
) -> None:
# node_ids is segment's node_ids
# Note: Summary indexes are now disabled (not deleted) when segments are disabled.
# This method is called for actual deletion scenarios (e.g., when segment is deleted).
@@ -160,24 +167,23 @@ class ParentChildIndexProcessor(BaseIndexProcessor):
if delete_summaries:
if node_ids:
# Find segments by index_node_id
with session_factory.create_session() as session:
segments = session.scalars(
select(DocumentSegment).where(
DocumentSegment.dataset_id == dataset.id,
DocumentSegment.index_node_id.in_(node_ids),
)
).all()
segment_ids = [segment.id for segment in segments]
if segment_ids:
SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=segment_ids)
segments = session.scalars(
select(DocumentSegment).where(
DocumentSegment.dataset_id == dataset.id,
DocumentSegment.index_node_id.in_(node_ids),
)
).all()
segment_ids = [segment.id for segment in segments]
if segment_ids:
SummaryIndexService.delete_summaries_for_segments(dataset, segment_ids, session=session)
else:
# Delete all summaries for the dataset
SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=None)
SummaryIndexService.delete_summaries_for_segments(dataset, None, session=session)
if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
delete_child_chunks = kwargs.get("delete_child_chunks") or False
precomputed_child_node_ids = kwargs.get("precomputed_child_node_ids")
vector = Vector(dataset)
vector = Vector(dataset, session=session)
if node_ids:
# Use precomputed child_node_ids if available (to avoid race conditions)
@@ -185,7 +191,7 @@ class ParentChildIndexProcessor(BaseIndexProcessor):
child_node_ids = precomputed_child_node_ids
else:
# Fallback to original query (may fail if segments are already deleted)
rows = db.session.execute(
rows = session.execute(
select(ChildChunk.index_node_id)
.join(DocumentSegment, ChildChunk.segment_id == DocumentSegment.id)
.where(
@@ -202,23 +208,23 @@ class ParentChildIndexProcessor(BaseIndexProcessor):
# Delete from database
if delete_child_chunks and child_node_ids:
db.session.execute(
session.execute(
delete(ChildChunk).where(
ChildChunk.dataset_id == dataset.id, ChildChunk.index_node_id.in_(child_node_ids)
)
)
db.session.commit()
session.flush()
else:
vector.delete()
if delete_child_chunks:
# Use existing compound index: (tenant_id, dataset_id, ...)
db.session.execute(
session.execute(
delete(ChildChunk).where(
ChildChunk.tenant_id == dataset.tenant_id, ChildChunk.dataset_id == dataset.id
)
)
db.session.commit()
session.flush()
def _split_child_nodes(
self,
@@ -257,7 +263,7 @@ class ParentChildIndexProcessor(BaseIndexProcessor):
return child_nodes
@override
def index(self, dataset: Dataset, document: DatasetDocument, chunks: Any) -> None:
def index(self, dataset: Dataset, document: DatasetDocument, chunks: Any, session: Session) -> None:
parent_childs = ParentChildStructureChunk.model_validate(chunks)
documents = []
for parent_child in parent_childs.parent_child_chunks:
@@ -291,10 +297,11 @@ class ParentChildIndexProcessor(BaseIndexProcessor):
attachments.append(file_document)
doc.attachments = attachments
else:
account = AccountService.load_user(document.created_by, db.session())
with session_factory.create_session() as account_session:
account = AccountService.load_user(document.created_by, account_session)
if not account:
raise ValueError("Invalid account")
doc.attachments = self._get_content_files(doc, current_user=account)
doc.attachments = self._get_content_files(doc, current_user=account, session=session)
documents.append(doc)
if documents:
# update document parent mode
@@ -308,14 +315,14 @@ class ParentChildIndexProcessor(BaseIndexProcessor):
),
created_by=document.created_by,
)
db.session.add(dataset_process_rule)
db.session.flush()
session.add(dataset_process_rule)
session.flush()
document.dataset_process_rule_id = dataset_process_rule.id
db.session.commit()
# save node to document segment
doc_store = DatasetDocumentStore(dataset=dataset, user_id=document.created_by, document_id=document.id)
# add document segments
doc_store.add_documents(docs=documents, save_child=True)
doc_store.add_documents(docs=documents, save_child=True, session=session)
session.commit()
if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
all_child_documents = []
all_multimodal_documents = []
@@ -324,7 +331,7 @@ class ParentChildIndexProcessor(BaseIndexProcessor):
all_child_documents.extend(doc.children)
if doc.attachments:
all_multimodal_documents.extend(doc.attachments)
vector = Vector(dataset)
vector = Vector(dataset, session=session)
if all_child_documents:
vector.create(all_child_documents)
if all_multimodal_documents and dataset.is_multimodal:
@@ -351,6 +358,8 @@ class ParentChildIndexProcessor(BaseIndexProcessor):
preview_texts: list[PreviewDetail],
summary_index_setting: SummaryIndexSettingDict,
doc_language: str | None = None,
*,
session: Session,
) -> list[PreviewDetail]:
"""
For each parent chunk in preview_texts, concurrently call generate_summary to generate a summary
@@ -377,21 +386,25 @@ class ParentChildIndexProcessor(BaseIndexProcessor):
if flask_app:
# Ensure Flask app context in worker thread
with flask_app.app_context():
with session_factory.create_session() as worker_session:
summary, _ = ParagraphIndexProcessor.generate_summary(
tenant_id=tenant_id,
text=preview.content,
summary_index_setting=summary_index_setting,
document_language=doc_language,
session=worker_session,
)
preview.summary = summary
else:
# Fallback: try without app context (may fail)
with session_factory.create_session() as worker_session:
summary, _ = ParagraphIndexProcessor.generate_summary(
tenant_id=tenant_id,
text=preview.content,
summary_index_setting=summary_index_setting,
document_language=doc_language,
session=worker_session,
)
preview.summary = summary
else:
# Fallback: try without app context (may fail)
summary, _ = ParagraphIndexProcessor.generate_summary(
tenant_id=tenant_id,
text=preview.content,
summary_index_setting=summary_index_setting,
document_language=doc_language,
)
preview.summary = summary
# Generate summaries concurrently using ThreadPoolExecutor
@@ -9,9 +9,9 @@ from typing import Any, TypedDict, override
import pandas as pd
from flask import Flask, current_app
from sqlalchemy import select
from sqlalchemy.orm import Session
from werkzeug.datastructures import FileStorage
from core.db.session_factory import session_factory
from core.entities.knowledge_entities import PreviewDetail
from core.llm_generator.llm_generator import LLMGenerator
from core.rag.cleaner.clean_processor import CleanProcessor
@@ -41,17 +41,20 @@ class QAFormatPreviewDict(TypedDict):
class QAIndexProcessor(BaseIndexProcessor):
@override
def extract(self, extract_setting: ExtractSetting, **kwargs) -> list[Document]:
def extract(self, extract_setting: ExtractSetting, *, session: Session, **kwargs) -> list[Document]:
text_docs = ExtractProcessor.extract(
extract_setting=extract_setting,
is_automatic=(
kwargs.get("process_rule_mode") == "automatic" or kwargs.get("process_rule_mode") == "hierarchical"
),
session=session,
)
return text_docs
@override
def transform(self, documents: list[Document], current_user: Account | None = None, **kwargs) -> list[Document]:
def transform(
self, documents: list[Document], current_user: Account | None = None, *, session: Session, **kwargs
) -> list[Document]:
preview = kwargs.get("preview")
process_rule = kwargs.get("process_rule")
if not process_rule:
@@ -145,16 +148,20 @@ class QAIndexProcessor(BaseIndexProcessor):
documents: list[Document],
multimodal_documents: list[AttachmentDocument] | None = None,
with_keywords: bool = True,
*,
session: Session,
**kwargs,
) -> None:
if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
vector = Vector(dataset)
vector = Vector(dataset, session=session)
vector.create(documents)
if multimodal_documents and dataset.is_multimodal:
vector.create_multimodal(multimodal_documents)
@override
def clean(self, dataset: Dataset, node_ids: list[str] | None, with_keywords: bool = True, **kwargs) -> None:
def clean(
self, dataset: Dataset, node_ids: list[str] | None, with_keywords: bool = True, *, session: Session, **kwargs
) -> None:
# Note: Summary indexes are now disabled (not deleted) when segments are disabled.
# This method is called for actual deletion scenarios (e.g., when segment is deleted).
# For disable operations, disable_summaries_for_segments is called directly in the task.
@@ -164,28 +171,27 @@ class QAIndexProcessor(BaseIndexProcessor):
if delete_summaries:
if node_ids:
# Find segments by index_node_id
with session_factory.create_session() as session:
segments = session.scalars(
select(DocumentSegment).where(
DocumentSegment.dataset_id == dataset.id,
DocumentSegment.index_node_id.in_(node_ids),
)
).all()
segment_ids = [segment.id for segment in segments]
if segment_ids:
SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=segment_ids)
segments = session.scalars(
select(DocumentSegment).where(
DocumentSegment.dataset_id == dataset.id,
DocumentSegment.index_node_id.in_(node_ids),
)
).all()
segment_ids = [segment.id for segment in segments]
if segment_ids:
SummaryIndexService.delete_summaries_for_segments(dataset, segment_ids, session=session)
else:
# Delete all summaries for the dataset
SummaryIndexService.delete_summaries_for_segments(dataset=dataset, segment_ids=None)
SummaryIndexService.delete_summaries_for_segments(dataset, None, session=session)
vector = Vector(dataset)
vector = Vector(dataset, session=session)
if node_ids:
vector.delete_by_ids(node_ids)
else:
vector.delete()
@override
def index(self, dataset: Dataset, document: DatasetDocument, chunks: Any) -> None:
def index(self, dataset: Dataset, document: DatasetDocument, chunks: Any, session: Session) -> None:
qa_chunks = QAStructureChunk.model_validate(chunks)
documents = []
for qa_chunk in qa_chunks.qa_chunks:
@@ -201,9 +207,10 @@ class QAIndexProcessor(BaseIndexProcessor):
if documents:
# save node to document segment
doc_store = DatasetDocumentStore(dataset=dataset, user_id=document.created_by, document_id=document.id)
doc_store.add_documents(docs=documents, save_child=False)
doc_store.add_documents(docs=documents, save_child=False, session=session)
session.commit()
if dataset.indexing_technique == IndexTechniqueType.HIGH_QUALITY:
vector = Vector(dataset)
vector = Vector(dataset, session=session)
vector.create(documents)
else:
raise ValueError("Indexing technique must be high quality.")
@@ -228,6 +235,8 @@ class QAIndexProcessor(BaseIndexProcessor):
preview_texts: list[PreviewDetail],
summary_index_setting: SummaryIndexSettingDict,
doc_language: str | None = None,
*,
session: Session,
) -> list[PreviewDetail]:
"""
QA model doesn't generate summaries, so this method returns preview_texts unchanged.
+8 -6
View File
@@ -1,12 +1,13 @@
import base64
from typing import override
from sqlalchemy.orm import Session
from core.model_manager import ModelInstance, ModelManager
from core.rag.index_processor.constant.doc_type import DocType
from core.rag.index_processor.constant.query_type import QueryType
from core.rag.models.document import Document
from core.rag.rerank.rerank_base import BaseRerankRunner
from extensions.ext_database import db
from extensions.ext_storage import storage
from graphon.model_runtime.entities.model_entities import ModelType
from graphon.model_runtime.entities.rerank_entities import MultimodalRerankInput, RerankResult
@@ -14,8 +15,11 @@ from models.model import UploadFile
class RerankModelRunner(BaseRerankRunner):
def __init__(self, rerank_model_instance: ModelInstance):
_session: Session
def __init__(self, rerank_model_instance: ModelInstance, *, session: Session):
self.rerank_model_instance = rerank_model_instance
self._session = session
@override
def run(
@@ -134,8 +138,7 @@ class RerankModelRunner(BaseRerankRunner):
and document.metadata["doc_id"] not in doc_ids
):
if document.metadata.get("doc_type") == DocType.IMAGE:
# Query file info within db.session context to ensure thread-safe access
upload_file = db.session.get(UploadFile, document.metadata["doc_id"])
upload_file = self._session.get(UploadFile, document.metadata["doc_id"])
if upload_file:
blob = storage.load_once(upload_file.key)
document_file_base64 = base64.b64encode(blob).decode()
@@ -169,8 +172,7 @@ class RerankModelRunner(BaseRerankRunner):
rerank_result, unique_documents = self.fetch_text_rerank(query, documents, score_threshold, top_n)
return rerank_result, unique_documents
elif query_type == QueryType.IMAGE_QUERY:
# Query file info within db.session context to ensure thread-safe access
upload_file = db.session.get(UploadFile, query)
upload_file = self._session.get(UploadFile, query)
if upload_file:
blob = storage.load_once(upload_file.key)
file_query = base64.b64encode(blob).decode()
+33 -19
View File
@@ -274,7 +274,8 @@ class DatasetRetrieval:
retrieval_resource_list.append(source)
# deal with dify documents
if dify_documents:
records = RetrievalService.format_retrieval_documents(dify_documents)
with Session(bind=session.get_bind()) as format_session:
records = RetrievalService.format_retrieval_documents(format_session, dify_documents)
dataset_ids = [i.segment.dataset_id for i in records]
document_ids = [i.segment.document_id for i in records]
@@ -491,7 +492,8 @@ class DatasetRetrieval:
retrieval_resource_list.append(source)
# deal with dify documents
if dify_documents:
records = RetrievalService.format_retrieval_documents(dify_documents)
with Session(bind=session.get_bind()) as format_session:
records = RetrievalService.format_retrieval_documents(format_session, dify_documents)
if records:
for record in records:
segment = record.segment
@@ -1225,7 +1227,11 @@ class DatasetRetrieval:
continue
# pass if dataset is not available
if dataset and dataset.provider != "external" and dataset.available_document_count == 0:
if (
dataset
and dataset.provider != "external"
and dataset.get_total_available_documents(session=session) == 0
):
continue
available_datasets.append(dataset)
@@ -1859,23 +1865,31 @@ class DatasetRetrieval:
# Skip second reranking when there is only one dataset
if reranking_enable and dataset_count > 1:
# do rerank for searched documents
data_post_processor = DataPostProcessor(tenant_id, reranking_mode, reranking_model, weights, False)
if query:
all_documents_item = data_post_processor.invoke(
query=query,
documents=all_documents_item,
score_threshold=score_threshold,
top_n=top_k,
query_type=QueryType.TEXT_QUERY,
)
if attachment_id:
all_documents_item = data_post_processor.invoke(
documents=all_documents_item,
score_threshold=score_threshold,
top_n=top_k,
query_type=QueryType.IMAGE_QUERY,
query=attachment_id,
with session_factory.create_session() as session:
data_post_processor = DataPostProcessor(
tenant_id,
reranking_mode,
reranking_model,
weights,
False,
session=session,
)
if query:
all_documents_item = data_post_processor.invoke(
query=query,
documents=all_documents_item,
score_threshold=score_threshold,
top_n=top_k,
query_type=QueryType.TEXT_QUERY,
)
if attachment_id:
all_documents_item = data_post_processor.invoke(
documents=all_documents_item,
score_threshold=score_threshold,
top_n=top_k,
query_type=QueryType.IMAGE_QUERY,
query=attachment_id,
)
else:
if index_type == IndexTechniqueType.ECONOMY:
if not query:
@@ -76,11 +76,11 @@ class DatasetMultiRetrieverTool(DatasetRetrieverBaseTool):
model=self.reranking_model_name,
)
rerank_runner = RerankModelRunner(rerank_model_instance)
rerank_runner = RerankModelRunner(rerank_model_instance, session=session)
all_documents = rerank_runner.run(query, all_documents, self.score_threshold, self.top_k)
for hit_callback in self.hit_callbacks:
hit_callback.on_tool_end(all_documents, db.session())
hit_callback.on_tool_end(all_documents, session)
document_score_list = {}
for item in all_documents:
@@ -96,7 +96,7 @@ class DatasetMultiRetrieverTool(DatasetRetrieverBaseTool):
DocumentSegment.enabled == True,
DocumentSegment.index_node_id.in_(index_node_ids),
)
segments = db.session.scalars(document_segment_stmt).all()
segments = session.scalars(document_segment_stmt).all()
if segments:
index_node_id_to_position = {id: position for position, id in enumerate(index_node_ids)}
@@ -112,13 +112,13 @@ class DatasetMultiRetrieverTool(DatasetRetrieverBaseTool):
context_list: list[RetrievalSourceMetadata] = []
resource_number = 1
for segment in sorted_segments:
dataset = db.session.get(Dataset, segment.dataset_id)
dataset = session.get(Dataset, segment.dataset_id)
document_stmt = select(Document).where(
Document.id == segment.document_id,
Document.enabled == True,
Document.archived == False,
)
document = db.session.scalar(document_stmt)
document = session.scalar(document_stmt)
if dataset and document:
source = RetrievalSourceMetadata(
position=resource_number,
@@ -12,7 +12,6 @@ from core.rag.models.document import Document as RetrievalDocument
from core.rag.retrieval.dataset_retrieval import DatasetRetrieval
from core.rag.retrieval.retrieval_methods import RetrievalMethod
from core.tools.utils.dataset_retriever.dataset_retriever_base_tool import DatasetRetrieverBaseTool
from extensions.ext_database import db
from models.dataset import Dataset
from models.dataset import Document as DatasetDocument
from services.external_knowledge_service import ExternalDatasetService
@@ -60,12 +59,12 @@ class DatasetRetrieverTool(DatasetRetrieverBaseTool):
@override
def _run(self, session: Session, query: str) -> str:
dataset_stmt = select(Dataset).where(Dataset.tenant_id == self.tenant_id, Dataset.id == self.dataset_id)
dataset = db.session.scalar(dataset_stmt)
dataset = session.scalar(dataset_stmt)
if not dataset:
return ""
for hit_callback in self.hit_callbacks:
hit_callback.on_query(query, dataset.id, db.session())
hit_callback.on_query(query, dataset.id, session)
dataset_retrieval = DatasetRetrieval()
metadata_filter_document_ids, metadata_condition = dataset_retrieval.get_metadata_filter_condition(
session,
@@ -162,14 +161,15 @@ class DatasetRetrieverTool(DatasetRetrieverBaseTool):
else:
documents = []
for hit_callback in self.hit_callbacks:
hit_callback.on_tool_end(documents, db.session())
hit_callback.on_tool_end(documents, session)
document_score_list = {}
if dataset.indexing_technique != IndexTechniqueType.ECONOMY:
for item in documents:
if item.metadata is not None and item.metadata.get("score"):
document_score_list[item.metadata["doc_id"]] = item.metadata["score"]
document_context_list: list[DocumentContext] = []
records = RetrievalService.format_retrieval_documents(documents)
with Session(bind=session.get_bind()) as format_session:
records = RetrievalService.format_retrieval_documents(format_session, documents)
if records:
for record in records:
segment = record.segment
@@ -195,13 +195,13 @@ class DatasetRetrieverTool(DatasetRetrieverBaseTool):
if self.return_resource:
for record in records:
segment = record.segment
dataset = db.session.get(Dataset, segment.dataset_id)
dataset = session.get(Dataset, segment.dataset_id)
dataset_document_stmt = select(DatasetDocument).where(
DatasetDocument.id == segment.document_id,
DatasetDocument.enabled == True,
DatasetDocument.archived == False,
)
document = db.session.scalar(dataset_document_stmt)
document = session.scalar(dataset_document_stmt)
if dataset and document:
source = RetrievalSourceMetadata(
dataset_id=dataset.id,
+1 -1
View File
@@ -281,7 +281,7 @@ class WorkflowTool(Tool):
user_stmt = select(Account).where(Account.id == user_id)
user = session.scalar(user_stmt)
if user:
user.current_tenant = tenant
user.set_current_tenant_with_session(tenant, session=session)
session.expunge(user)
return user
@@ -147,7 +147,7 @@ class WorkflowAgentNodeValidator:
)
cls._validate_agent_soul_env(binding=binding, agent_soul=agent_soul)
cls._validate_agent_soul_tools(binding=binding, agent_soul=agent_soul)
cls._validate_agent_soul_knowledge(binding=binding, agent_soul=agent_soul)
cls._validate_agent_soul_knowledge(session=session, binding=binding, agent_soul=agent_soul)
node_job = WorkflowNodeJobConfig.model_validate(binding.node_job_config_dict)
cls.validate_node_job(session=session, binding=binding, node_job=node_job, topology=topology)
@@ -370,11 +370,13 @@ class WorkflowAgentNodeValidator:
def _validate_agent_soul_knowledge(
cls,
*,
session: Session,
binding: WorkflowAgentNodeBinding,
agent_soul: AgentSoulConfig,
) -> None:
"""Validate knowledge set dataset rows against the publishing tenant."""
missing_ids = list_missing_tenant_knowledge_dataset_ids(
session=session,
tenant_id=binding.tenant_id,
agent_soul=agent_soul,
)
@@ -2,6 +2,9 @@ import logging
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, override
from sqlalchemy.orm import Session
from core.db.session_factory import session_factory
from core.rag.index_processor.index_processor import IndexProcessor
from core.rag.index_processor.index_processor_base import SummaryIndexSettingDict
from core.rag.summary_index.summary_index import SummaryIndex
@@ -83,9 +86,15 @@ class KnowledgeIndexNode(Node[KnowledgeIndexNodeData]):
# Get indexing_technique and summary_index_setting from node_data (workflow graph config)
# or fallback to dataset if not available in node_data
outputs = self.index_processor.get_preview_output(
chunks, dataset_id, document_id, node_data.chunk_structure, summary_index_setting
)
with session_factory.create_session() as session:
outputs = self.index_processor.get_preview_output(
chunks,
dataset_id,
document_id,
node_data.chunk_structure,
summary_index_setting,
session=session,
)
return NodeRunResult(
status=WorkflowNodeExecutionStatus.SUCCEEDED,
inputs=variables,
@@ -97,15 +106,17 @@ class KnowledgeIndexNode(Node[KnowledgeIndexNodeData]):
if not batch:
raise KnowledgeIndexNodeError("Batch is required.")
results = self._invoke_knowledge_index(
dataset_id=dataset_id,
document_id=document_id,
original_document_id=original_document_id_segment.value if original_document_id_segment else "",
is_preview=is_preview,
batch=batch.value,
chunks=chunks,
summary_index_setting=summary_index_setting,
)
with session_factory.create_session() as session:
results = self._invoke_knowledge_index(
session=session,
dataset_id=dataset_id,
document_id=document_id,
original_document_id=original_document_id_segment.value if original_document_id_segment else "",
is_preview=is_preview,
batch=batch.value,
chunks=chunks,
summary_index_setting=summary_index_setting,
)
return NodeRunResult(status=WorkflowNodeExecutionStatus.SUCCEEDED, inputs=variables, outputs=results)
except KnowledgeIndexNodeError as e:
@@ -134,12 +145,16 @@ class KnowledgeIndexNode(Node[KnowledgeIndexNodeData]):
batch: Any,
chunks: Mapping[str, Any],
summary_index_setting: SummaryIndexSettingDict | None = None,
*,
session: Session,
):
if not document_id:
raise KnowledgeIndexNodeError("document_id is required.")
rst = self.index_processor.index_and_clean(
dataset_id, document_id, original_document_id, chunks, batch, summary_index_setting
dataset_id, document_id, original_document_id, chunks, batch, summary_index_setting, session=session
)
# Summary generation opens independent sessions and must see the indexed rows.
session.commit()
self.summary_index_service.generate_and_vectorize_summary(
dataset_id, document_id, is_preview, summary_index_setting
)
@@ -2,6 +2,7 @@ from collections.abc import Mapping
from typing import Any, Protocol, TypedDict
from pydantic import BaseModel, Field
from sqlalchemy.orm import Session
class IndexingResultDict(TypedDict):
@@ -44,6 +45,8 @@ class IndexProcessorProtocol(Protocol):
chunks: Mapping[str, Any],
batch: Any,
summary_index_setting: dict[str, Any] | None = None,
*,
session: Session,
) -> IndexingResultDict: ...
def get_preview_output(
@@ -53,6 +56,8 @@ class IndexProcessorProtocol(Protocol):
document_id: str,
chunk_structure: str,
summary_index_setting: dict[str, Any] | None,
*,
session: Session,
) -> Preview: ...