refactor(channels): 简化sqlalchemy布尔值转换写法

移除冗余的func.cast嵌套调用,直接使用cast方法转换布尔值类型,代码更简洁易读
This commit is contained in:
Kris 2026-05-13 18:05:12 +08:00
parent 29cf49ec92
commit 0f89ccf90a
2 changed files with 9 additions and 30 deletions

View File

@ -5,7 +5,7 @@ import time
from collections import defaultdict, deque from collections import defaultdict, deque
from typing import Any from typing import Any
from sqlalchemy import func, select from sqlalchemy import Integer, func, select
from yuxi.channels.base import BaseChannelAdapter from yuxi.channels.base import BaseChannelAdapter
from yuxi.channels.exceptions import ChannelException from yuxi.channels.exceptions import ChannelException
@ -655,16 +655,10 @@ class ChannelManager:
ChannelMsgRecord.channel_id, ChannelMsgRecord.channel_id,
func.count().label("total"), func.count().label("total"),
func.sum( func.sum(
func.cast( (ChannelMsgRecord.status == "success").cast(Integer),
(ChannelMsgRecord.status == "success").cast(func.Integer),
func.Integer,
)
).label("success_count"), ).label("success_count"),
func.sum( func.sum(
func.cast( (ChannelMsgRecord.status == "error").cast(Integer),
(ChannelMsgRecord.status == "error").cast(func.Integer),
func.Integer,
)
).label("error_count"), ).label("error_count"),
) )
.where(ChannelMsgRecord.channel_id.in_(channel_ids)) .where(ChannelMsgRecord.channel_id.in_(channel_ids))
@ -798,16 +792,10 @@ class ChannelManager:
select( select(
func.count().label("total"), func.count().label("total"),
func.sum( func.sum(
func.cast( (ChannelMsgRecord.status == "success").cast(Integer),
(ChannelMsgRecord.status == "success").cast(func.Integer),
func.Integer,
)
).label("success_count"), ).label("success_count"),
func.sum( func.sum(
func.cast( (ChannelMsgRecord.status == "error").cast(Integer),
(ChannelMsgRecord.status == "error").cast(func.Integer),
func.Integer,
)
).label("error_count"), ).label("error_count"),
).where(ChannelMsgRecord.channel_id == channel_id) ).where(ChannelMsgRecord.channel_id == channel_id)
) )

View File

@ -240,7 +240,7 @@ async def get_channel_stats(
): ):
from datetime import timedelta from datetime import timedelta
from sqlalchemy import func, select from sqlalchemy import Integer, func
from yuxi.storage.postgres.models_channels import ChannelMsgRecord from yuxi.storage.postgres.models_channels import ChannelMsgRecord
from yuxi.utils.datetime_utils import utc_now_naive from yuxi.utils.datetime_utils import utc_now_naive
@ -260,16 +260,10 @@ async def get_channel_stats(
stats_base = select( stats_base = select(
func.count().label("total"), func.count().label("total"),
func.sum( func.sum(
func.cast( (ChannelMsgRecord.status == "success").cast(Integer),
(ChannelMsgRecord.status == "success").cast(func.Integer),
func.Integer,
)
).label("success_count"), ).label("success_count"),
func.sum( func.sum(
func.cast( (ChannelMsgRecord.status == "error").cast(Integer),
(ChannelMsgRecord.status == "error").cast(func.Integer),
func.Integer,
)
).label("error_count"), ).label("error_count"),
func.avg(ChannelMsgRecord.response_time_ms).label("avg_response_ms"), func.avg(ChannelMsgRecord.response_time_ms).label("avg_response_ms"),
func.percentile_cont(0.5).within_group(ChannelMsgRecord.response_time_ms).label("p50_ms"), func.percentile_cont(0.5).within_group(ChannelMsgRecord.response_time_ms).label("p50_ms"),
@ -343,10 +337,7 @@ async def get_channel_stats(
func.date(ChannelMsgRecord.created_at).label("date"), func.date(ChannelMsgRecord.created_at).label("date"),
func.count().label("total"), func.count().label("total"),
func.sum( func.sum(
func.cast( (ChannelMsgRecord.status == "success").cast(Integer),
(ChannelMsgRecord.status == "success").cast(func.Integer),
func.Integer,
)
).label("success"), ).label("success"),
) )
.where( .where(