refactor: select in message_service and ops_service (#34414)

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
Renzo
2026-04-01 16:37:27 +00:00
committed by GitHub
co-authored by autofix-ci[bot]
parent 391007d02e
commit 4e1d060439
4 changed files with 94 additions and 195 deletions
+26 -31
View File
@@ -3,6 +3,7 @@ from typing import Union
from graphon.model_runtime.entities.model_entities import ModelType
from pydantic import TypeAdapter
from sqlalchemy import select
from sqlalchemy.orm import sessionmaker
from core.app.apps.advanced_chat.app_config_manager import AdvancedChatAppConfigManager
@@ -75,17 +76,15 @@ class MessageService:
fetch_limit = limit + 1
if first_id:
first_message = (
db.session.query(Message)
.where(Message.conversation_id == conversation.id, Message.id == first_id)
.first()
first_message = db.session.scalar(
select(Message).where(Message.conversation_id == conversation.id, Message.id == first_id).limit(1)
)
if not first_message:
raise FirstMessageNotExistsError()
history_messages = (
db.session.query(Message)
history_messages = db.session.scalars(
select(Message)
.where(
Message.conversation_id == conversation.id,
Message.created_at < first_message.created_at,
@@ -93,16 +92,14 @@ class MessageService:
)
.order_by(Message.created_at.desc())
.limit(fetch_limit)
.all()
)
).all()
else:
history_messages = (
db.session.query(Message)
history_messages = db.session.scalars(
select(Message)
.where(Message.conversation_id == conversation.id)
.order_by(Message.created_at.desc())
.limit(fetch_limit)
.all()
)
).all()
has_more = False
if len(history_messages) > limit:
@@ -129,7 +126,7 @@ class MessageService:
if not user:
return InfiniteScrollPagination(data=[], limit=limit, has_more=False)
base_query = db.session.query(Message)
stmt = select(Message)
fetch_limit = limit + 1
@@ -138,28 +135,27 @@ class MessageService:
app_model=app_model, user=user, conversation_id=conversation_id
)
base_query = base_query.where(Message.conversation_id == conversation.id)
stmt = stmt.where(Message.conversation_id == conversation.id)
# Check if include_ids is not None and not empty to avoid WHERE false condition
if include_ids is not None:
if len(include_ids) == 0:
return InfiniteScrollPagination(data=[], limit=limit, has_more=False)
base_query = base_query.where(Message.id.in_(include_ids))
stmt = stmt.where(Message.id.in_(include_ids))
if last_id:
last_message = base_query.where(Message.id == last_id).first()
last_message = db.session.scalar(stmt.where(Message.id == last_id).limit(1))
if not last_message:
raise LastMessageNotExistsError()
history_messages = (
base_query.where(Message.created_at < last_message.created_at, Message.id != last_message.id)
history_messages = db.session.scalars(
stmt.where(Message.created_at < last_message.created_at, Message.id != last_message.id)
.order_by(Message.created_at.desc())
.limit(fetch_limit)
.all()
)
).all()
else:
history_messages = base_query.order_by(Message.created_at.desc()).limit(fetch_limit).all()
history_messages = db.session.scalars(stmt.order_by(Message.created_at.desc()).limit(fetch_limit)).all()
has_more = False
if len(history_messages) > limit:
@@ -214,21 +210,20 @@ class MessageService:
def get_all_messages_feedbacks(cls, app_model: App, page: int, limit: int):
"""Get all feedbacks of an app"""
offset = (page - 1) * limit
feedbacks = (
db.session.query(MessageFeedback)
feedbacks = db.session.scalars(
select(MessageFeedback)
.where(MessageFeedback.app_id == app_model.id)
.order_by(MessageFeedback.created_at.desc(), MessageFeedback.id.desc())
.limit(limit)
.offset(offset)
.all()
)
).all()
return [record.to_dict() for record in feedbacks]
@classmethod
def get_message(cls, app_model: App, user: Union[Account, EndUser] | None, message_id: str):
message = (
db.session.query(Message)
message = db.session.scalar(
select(Message)
.where(
Message.id == message_id,
Message.app_id == app_model.id,
@@ -236,7 +231,7 @@ class MessageService:
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),
)
.first()
.limit(1)
)
if not message:
@@ -282,10 +277,10 @@ class MessageService:
)
else:
if not conversation.override_model_configs:
app_model_config = (
db.session.query(AppModelConfig)
app_model_config = db.session.scalar(
select(AppModelConfig)
.where(AppModelConfig.id == conversation.app_model_config_id, AppModelConfig.app_id == app_model.id)
.first()
.limit(1)
)
else:
conversation_override_model_configs = _app_model_config_adapter.validate_json(
+17 -15
View File
@@ -1,5 +1,7 @@
from typing import Any
from sqlalchemy import select
from core.ops.entities.config_entity import BaseTracingConfig
from core.ops.ops_trace_manager import OpsTraceManager, provider_config_map
from extensions.ext_database import db
@@ -15,17 +17,17 @@ class OpsService:
:param tracing_provider: tracing provider
:return:
"""
trace_config_data: TraceAppConfig | None = (
db.session.query(TraceAppConfig)
trace_config_data: TraceAppConfig | None = db.session.scalar(
select(TraceAppConfig)
.where(TraceAppConfig.app_id == app_id, TraceAppConfig.tracing_provider == tracing_provider)
.first()
.limit(1)
)
if not trace_config_data:
return None
# decrypt_token and obfuscated_token
app = db.session.query(App).where(App.id == app_id).first()
app = db.session.get(App, app_id)
if not app:
return None
tenant_id = app.tenant_id
@@ -182,17 +184,17 @@ class OpsService:
project_url = None
# check if trace config already exists
trace_config_data: TraceAppConfig | None = (
db.session.query(TraceAppConfig)
trace_config_data: TraceAppConfig | None = db.session.scalar(
select(TraceAppConfig)
.where(TraceAppConfig.app_id == app_id, TraceAppConfig.tracing_provider == tracing_provider)
.first()
.limit(1)
)
if trace_config_data:
return None
# get tenant id
app = db.session.query(App).where(App.id == app_id).first()
app = db.session.get(App, app_id)
if not app:
return None
tenant_id = app.tenant_id
@@ -224,17 +226,17 @@ class OpsService:
raise ValueError(f"Invalid tracing provider: {tracing_provider}")
# check if trace config already exists
current_trace_config = (
db.session.query(TraceAppConfig)
current_trace_config = db.session.scalar(
select(TraceAppConfig)
.where(TraceAppConfig.app_id == app_id, TraceAppConfig.tracing_provider == tracing_provider)
.first()
.limit(1)
)
if not current_trace_config:
return None
# get tenant id
app = db.session.query(App).where(App.id == app_id).first()
app = db.session.get(App, app_id)
if not app:
return None
tenant_id = app.tenant_id
@@ -261,10 +263,10 @@ class OpsService:
:param tracing_provider: tracing provider
:return:
"""
trace_config = (
db.session.query(TraceAppConfig)
trace_config = db.session.scalar(
select(TraceAppConfig)
.where(TraceAppConfig.app_id == app_id, TraceAppConfig.tracing_provider == tracing_provider)
.first()
.limit(1)
)
if not trace_config: