feat: 缓存计费细分、能力匹配优化、用户模型调用计数

1. 缓存创建 tokens 区分 5min/1h TTL,支持按缓存时长差异化计费
   - Usage 表新增 cache_creation_input_tokens_5m/1h 字段
   - Claude handler 解析新格式 (ephemeral_5m/1h, claude_cache_creation_5/1h)
   - 计费规则支持 cache_ttl_pricing 覆盖 cache_creation 价格

2. 能力匹配机制优化
   - COMPATIBLE 能力不再硬过滤,改为排序阶段通过 capability_miss_count 优先级处理
   - cache_1h 改为 COMPATIBLE + REQUEST_PARAM(自动检测请求体中的 ttl=1h)
   - gemini_files 改为 EXCLUSIVE + REQUEST_PARAM(自动检测 fileData.fileUri)
   - 移除前端模型偏好/能力配置 UI(不再需要用户手动配置)

3. 新增用户-模型维度调用次数计数器 (UserModelUsageCount)
   - 原子递增,避免从 Usage 表聚合查询
   - 前端模型目录和用户可用模型列表展示调用次数

4. 其他改进
   - global_model_id 改为必填(NOT NULL),清理孤立模型
   - 模型映射对话框支持从上游获取模型列表并分组折叠
   - 端点测试不再依赖端点启用状态
   - 异步任务页面对普通用户隐藏用户信息列
   - Dashboard 响应式布局断点调整 (sm -> lg)
   - 号池管理仅展示已启用号池的提供商
This commit is contained in:
fawney19
2026-02-28 11:44:08 +08:00
parent 82bbed2720
commit ecb16d345a
59 changed files with 1053 additions and 608 deletions

View File

@@ -0,0 +1,200 @@
"""Add cache_creation columns, clean up capability settings, add user_model_usage_counts,
enforce global_model_id NOT NULL
1. Add cache_creation_input_tokens_5m and cache_creation_input_tokens_1h to usage table.
2. Clean up cache_1h/context_1m/gemini_files from user-configurable settings
(now auto-detected via REQUEST_PARAM mode).
3. Create user_model_usage_counts table for per-user per-model atomic usage counters.
4. Enforce models.global_model_id NOT NULL (delete orphan models without global model).
Revision ID: b2c3d4e5f6a7
Revises: 9a0b1c2d3e4f
Create Date: 2026-02-28 14:00:00.000000
"""
from __future__ import annotations
import json
import uuid
from collections.abc import Sequence
from datetime import datetime, timezone
import sqlalchemy as sa
from sqlalchemy import inspect
from alembic import op
revision: str = "b2c3d4e5f6a7"
down_revision: str | None = "9a0b1c2d3e4f"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def column_exists(table_name: str, column_name: str) -> bool:
bind = op.get_bind()
inspector = inspect(bind)
columns = [c["name"] for c in inspector.get_columns(table_name)]
return column_name in columns
def upgrade() -> None:
# --- 1. Add cache_creation columns ---
if not column_exists("usage", "cache_creation_input_tokens_5m"):
op.add_column(
"usage",
sa.Column(
"cache_creation_input_tokens_5m",
sa.Integer(),
nullable=False,
server_default=sa.text("0"),
comment="5min TTL cache creation input tokens",
),
)
if not column_exists("usage", "cache_creation_input_tokens_1h"):
op.add_column(
"usage",
sa.Column(
"cache_creation_input_tokens_1h",
sa.Integer(),
nullable=False,
server_default=sa.text("0"),
comment="1h TTL cache creation input tokens",
),
)
# --- 2. Clean up stale capability settings (pure Python, DB-agnostic) ---
stale_keys = {"cache_1h", "context_1m", "gemini_files"}
conn = op.get_bind()
# ApiKey.force_capabilities: dict-like JSON, remove stale keys
rows = conn.execute(
sa.text("SELECT id, force_capabilities FROM api_keys WHERE force_capabilities IS NOT NULL")
).fetchall()
for row in rows:
raw = row[1]
if raw is None:
continue
data = raw if isinstance(raw, dict) else json.loads(raw)
cleaned = {k: v for k, v in data.items() if k not in stale_keys}
new_val = json.dumps(cleaned) if cleaned else None
conn.execute(
sa.text("UPDATE api_keys SET force_capabilities = :val WHERE id = :id"),
{"val": new_val, "id": row[0]},
)
# User.model_capability_settings: nested dict {model_key: {cap: val}}, remove stale keys
rows = conn.execute(
sa.text(
"SELECT id, model_capability_settings FROM users"
" WHERE model_capability_settings IS NOT NULL"
)
).fetchall()
for row in rows:
raw = row[1]
if raw is None:
continue
data = raw if isinstance(raw, dict) else json.loads(raw)
cleaned = {}
for model_key, caps in data.items():
cap_cleaned = {k: v for k, v in caps.items() if k not in stale_keys}
if cap_cleaned:
cleaned[model_key] = cap_cleaned
new_val = json.dumps(cleaned) if cleaned else None
conn.execute(
sa.text("UPDATE users SET model_capability_settings = :val WHERE id = :id"),
{"val": new_val, "id": row[0]},
)
# GlobalModel.supported_capabilities: JSON array, remove stale entries
rows = conn.execute(
sa.text(
"SELECT id, supported_capabilities FROM global_models"
" WHERE supported_capabilities IS NOT NULL"
)
).fetchall()
for row in rows:
raw = row[1]
if raw is None:
continue
data = raw if isinstance(raw, list) else json.loads(raw)
cleaned = [c for c in data if c not in stale_keys]
new_val = json.dumps(cleaned) if cleaned else None
conn.execute(
sa.text("UPDATE global_models SET supported_capabilities = :val WHERE id = :id"),
{"val": new_val, "id": row[0]},
)
# --- 3. Create user_model_usage_counts table ---
op.create_table(
"user_model_usage_counts",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column(
"user_id",
sa.String(36),
sa.ForeignKey("users.id", ondelete="CASCADE"),
nullable=False,
),
sa.Column("model", sa.String(100), nullable=False),
sa.Column("usage_count", sa.Integer, nullable=False, server_default="0"),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.func.now(),
),
sa.Column(
"updated_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.func.now(),
),
sa.UniqueConstraint("user_id", "model", name="uq_user_model_usage_count"),
)
op.create_index("idx_user_model_usage_user", "user_model_usage_counts", ["user_id"])
op.create_index("idx_user_model_usage_model", "user_model_usage_counts", ["model"])
# Backfill from existing usage records
rows = conn.execute(
sa.text(
"SELECT user_id, model, COUNT(*) AS cnt FROM usage"
" WHERE user_id IS NOT NULL GROUP BY user_id, model"
)
).fetchall()
now = datetime.now(timezone.utc)
for row in rows:
conn.execute(
sa.text(
"INSERT INTO user_model_usage_counts"
" (id, user_id, model, usage_count, created_at, updated_at)"
" VALUES (:id, :user_id, :model, :cnt, :now, :now)"
),
{
"id": str(uuid.uuid4()),
"user_id": row[0],
"model": row[1],
"cnt": row[2],
"now": now,
},
)
# --- 4. Enforce models.global_model_id NOT NULL ---
op.execute("DELETE FROM models WHERE global_model_id IS NULL")
op.alter_column("models", "global_model_id", existing_type=sa.String(36), nullable=False)
def downgrade() -> None:
# Revert models.global_model_id to nullable
op.alter_column("models", "global_model_id", existing_type=sa.String(36), nullable=True)
# Drop user_model_usage_counts
op.drop_index("idx_user_model_usage_model", table_name="user_model_usage_counts")
op.drop_index("idx_user_model_usage_user", table_name="user_model_usage_counts")
op.drop_table("user_model_usage_counts")
# Drop cache_creation columns
if column_exists("usage", "cache_creation_input_tokens_1h"):
op.drop_column("usage", "cache_creation_input_tokens_1h")
if column_exists("usage", "cache_creation_input_tokens_5m"):
op.drop_column("usage", "cache_creation_input_tokens_5m")
# capability settings cleanup is not reversible

View File

@@ -26,7 +26,7 @@ export interface TieredPricingConfig {
export interface Model { export interface Model {
id: string id: string
provider_id: string provider_id: string
global_model_id?: string // 关联的 GlobalModel ID global_model_id: string // 关联的 GlobalModel ID
provider_model_name: string // Provider 侧的主模型名称 provider_model_name: string // Provider 侧的主模型名称
provider_model_mappings?: ProviderModelMapping[] | null // 模型名称映射列表(带优先级) provider_model_mappings?: ProviderModelMapping[] | null // 模型名称映射列表(带优先级)
config?: Record<string, unknown> | null // 额外配置(如 billing/video 等) config?: Record<string, unknown> | null // 额外配置(如 billing/video 等)
@@ -251,8 +251,8 @@ export interface UpstreamModel {
export interface ImportFromUpstreamSuccessItem { export interface ImportFromUpstreamSuccessItem {
model_id: string model_id: string
provider_model_id: string provider_model_id: string
global_model_id?: string // 可选,未关联时为空字符串 global_model_id: string
global_model_name?: string // 可选,未关联时为空字符串 global_model_name: string
created_global_model: boolean // 始终为 false不再自动创建 GlobalModel created_global_model: boolean // 始终为 false不再自动创建 GlobalModel
} }

View File

@@ -257,6 +257,7 @@ export const meApi = {
default_tiered_pricing: TieredPricingConfig | null default_tiered_pricing: TieredPricingConfig | null
supported_capabilities: string[] | null supported_capabilities: string[] | null
config: Record<string, unknown> | null config: Record<string, unknown> | null
usage_count: number
}> }>
total: number total: number
}> { }> {

View File

@@ -17,6 +17,8 @@ export interface PublicGlobalModel {
supported_capabilities: string[] | null supported_capabilities: string[] | null
// 模型配置JSON // 模型配置JSON
config: Record<string, unknown> | null config: Record<string, unknown> | null
// 调用次数
usage_count: number
} }
export interface PublicGlobalModelListResponse { export interface PublicGlobalModelListResponse {

View File

@@ -139,26 +139,6 @@
</div> </div>
</div> </div>
<!-- 模型偏好 -->
<div
v-if="model.supported_capabilities && model.supported_capabilities.length > 0"
class="space-y-3"
>
<h4 class="font-semibold text-sm">
模型偏好
</h4>
<div class="flex flex-wrap gap-2">
<Badge
v-for="cap in model.supported_capabilities"
:key="cap"
variant="outline"
class="text-xs"
>
{{ getCapabilityDisplayName(cap) }}
</Badge>
</div>
</div>
<!-- 默认定价 --> <!-- 默认定价 -->
<div class="space-y-3"> <div class="space-y-3">
<h4 class="font-semibold text-sm"> <h4 class="font-semibold text-sm">
@@ -475,12 +455,10 @@ import { getGlobalModelRoutingPreview } from '@/api/global-models'
// 使用外部类型定义 // 使用外部类型定义
import type { GlobalModelResponse } from '@/api/global-models' import type { GlobalModelResponse } from '@/api/global-models'
import type { TieredPricingConfig, PricingTier, ModelRoutingPreviewResponse } from '@/api/endpoints/types' import type { TieredPricingConfig, PricingTier, ModelRoutingPreviewResponse } from '@/api/endpoints/types'
import type { CapabilityDefinition } from '@/api/endpoints'
import type { RoutingProviderInfo } from '@/api/global-models' import type { RoutingProviderInfo } from '@/api/global-models'
const props = withDefaults(defineProps<Props>(), { const props = withDefaults(defineProps<Props>(), {
hasBlockingDialogOpen: false, hasBlockingDialogOpen: false,
capabilities: () => [],
}) })
const emit = defineEmits<{ const emit = defineEmits<{
'update:open': [value: boolean] 'update:open': [value: boolean]
@@ -500,7 +478,6 @@ interface Props {
model: GlobalModelResponse | null model: GlobalModelResponse | null
open: boolean open: boolean
hasBlockingDialogOpen?: boolean hasBlockingDialogOpen?: boolean
capabilities?: CapabilityDefinition[]
} }
// RoutingTab 引用 // RoutingTab 引用
@@ -571,12 +548,6 @@ defineExpose({
refreshRoutingData refreshRoutingData
}) })
// 根据能力名称获取显示名称
function getCapabilityDisplayName(capName: string): string {
const cap = props.capabilities?.find(c => c.name === capName)
return cap?.display_name || capName
}
// 检测是否有视频分辨率计费配置 // 检测是否有视频分辨率计费配置
const hasVideoPricing = computed(() => { const hasVideoPricing = computed(() => {
const priceByResolution = props.model?.config?.billing?.video?.price_per_second_by_resolution const priceByResolution = props.model?.config?.billing?.video?.price_per_second_by_resolution

View File

@@ -179,7 +179,6 @@ const searchQuery = ref('')
const existingGlobalModelIds = computed(() => { const existingGlobalModelIds = computed(() => {
return new Set( return new Set(
existingModels.value existingModels.value
.filter(m => m.global_model_id)
.map(m => m.global_model_id) .map(m => m.global_model_id)
) )
}) })

View File

@@ -63,6 +63,33 @@
> >
已选 {{ selectedNames.length }} 已选 {{ selectedNames.length }}
</span> </span>
<!-- 刷新上游模型按钮 -->
<button
v-if="upstreamModelsLoaded"
type="button"
class="p-2 hover:bg-muted rounded-md transition-colors shrink-0"
:disabled="fetchingUpstreamModels"
title="刷新上游模型"
@click="fetchUpstreamModels()"
>
<RefreshCw
class="w-4 h-4"
:class="{ 'animate-spin': fetchingUpstreamModels }"
/>
</button>
<button
v-else-if="!fetchingUpstreamModels"
type="button"
class="p-2 hover:bg-muted rounded-md transition-colors shrink-0"
title="从提供商获取模型"
@click="fetchUpstreamModels()"
>
<Zap class="w-4 h-4" />
</button>
<Loader2
v-else
class="w-4 h-4 animate-spin text-muted-foreground shrink-0"
/>
</div> </div>
<!-- 模型列表 --> <!-- 模型列表 -->
@@ -97,14 +124,22 @@
<!-- 自定义映射名称 --> <!-- 自定义映射名称 -->
<div v-if="customNames.length > 0"> <div v-if="customNames.length > 0">
<div <div
class="flex items-center justify-between px-3 py-2 bg-muted sticky top-0 z-20" class="flex items-center justify-between px-3 py-2 bg-muted sticky top-0 z-20 cursor-pointer hover:bg-muted/80 transition-colors"
@click="toggleGroupCollapse('custom')"
> >
<div class="flex items-center gap-2"> <div class="flex items-center gap-2">
<ChevronDown
class="w-4 h-4 transition-transform shrink-0"
:class="collapsedGroups.has('custom') ? '-rotate-90' : ''"
/>
<span class="text-xs font-medium">自定义模型</span> <span class="text-xs font-medium">自定义模型</span>
<span class="text-xs text-muted-foreground">({{ customNames.length }})</span> <span class="text-xs text-muted-foreground">({{ customNames.length }})</span>
</div> </div>
</div> </div>
<div class="space-y-1 p-2"> <div
v-show="!collapsedGroups.has('custom')"
class="space-y-1 p-2"
>
<div <div
v-for="name in sortedCustomNames" v-for="name in sortedCustomNames"
:key="name" :key="name"
@@ -125,6 +160,52 @@
</div> </div>
</div> </div>
<!-- 上游模型 -->
<template v-if="filteredUpstreamModels.length > 0">
<div
class="flex items-center justify-between px-3 py-2 bg-muted sticky top-0 z-20 cursor-pointer hover:bg-muted/80 transition-colors"
@click="toggleGroupCollapse('upstream')"
>
<div class="flex items-center gap-2">
<ChevronDown
class="w-4 h-4 transition-transform shrink-0"
:class="collapsedGroups.has('upstream') ? '-rotate-90' : ''"
/>
<span class="text-xs font-medium">上游模型</span>
<span class="text-xs text-muted-foreground">({{ upstreamModelNames.length }})</span>
</div>
<button
type="button"
class="text-xs text-primary hover:underline"
@click.stop="toggleAllUpstreamModels"
>
{{ isAllUpstreamModelsSelected ? '取消全选' : '全选' }}
</button>
</div>
<div
v-show="!collapsedGroups.has('upstream')"
class="space-y-1 p-2"
>
<div
v-for="name in filteredUpstreamModels"
:key="name"
class="flex items-center gap-2 px-2 py-1.5 rounded hover:bg-muted cursor-pointer"
@click="toggleName(name)"
>
<div
class="w-4 h-4 border rounded flex items-center justify-center shrink-0"
:class="selectedNames.includes(name) ? 'bg-primary border-primary' : ''"
>
<Check
v-if="selectedNames.includes(name)"
class="w-3 h-3 text-primary-foreground"
/>
</div>
<span class="text-sm font-mono truncate flex-1">{{ name }}</span>
</div>
</div>
</template>
<!-- 空状态 --> <!-- 空状态 -->
<div <div
v-if="showEmptyState" v-if="showEmptyState"
@@ -167,7 +248,7 @@
<script setup lang="ts"> <script setup lang="ts">
import { ref, computed, watch } from 'vue' import { ref, computed, watch } from 'vue'
import { Tag, Loader2, Plus, Search, Check } from 'lucide-vue-next' import { Tag, Loader2, Plus, Search, Check, ChevronDown, RefreshCw, Zap } from 'lucide-vue-next'
import { import {
Button, Button,
Input, Input,
@@ -184,8 +265,10 @@ import { parseApiError } from '@/utils/errorParser'
import { import {
type Model, type Model,
type ProviderModelAlias, type ProviderModelAlias,
type UpstreamModel,
} from '@/api/endpoints' } from '@/api/endpoints'
import { updateModel } from '@/api/endpoints/models' import { updateModel } from '@/api/endpoints/models'
import { useUpstreamModelsCache } from '../composables/useUpstreamModelsCache'
export interface AliasGroup { export interface AliasGroup {
model: Model model: Model
@@ -204,6 +287,7 @@ const props = defineProps<{
models: Model[] models: Model[]
editingGroup?: AliasGroup | null editingGroup?: AliasGroup | null
preselectedModelId?: string | null preselectedModelId?: string | null
hasAutoFetchKey?: boolean
}>() }>()
const emit = defineEmits<{ const emit = defineEmits<{
@@ -212,14 +296,23 @@ const emit = defineEmits<{
}>() }>()
const { error: showError, success: showSuccess } = useToast() const { error: showError, success: showSuccess } = useToast()
const { fetchModels: fetchCachedModels } = useUpstreamModelsCache()
// 状态 // 状态
const submitting = ref(false) const submitting = ref(false)
const loadingModels = ref(false) const loadingModels = ref(false)
const fetchingUpstreamModels = ref(false)
const upstreamModelsLoaded = ref(false)
// 搜索 // 搜索
const searchQuery = ref('') const searchQuery = ref('')
// 折叠状态
const collapsedGroups = ref<Set<string>>(new Set())
// 上游模型
const upstreamModels = ref<UpstreamModel[]>([])
// 表单数据 // 表单数据
const formData = ref<{ const formData = ref<{
modelId: string modelId: string
@@ -233,9 +326,26 @@ const selectedNames = ref<string[]>([])
// 自定义名称列表(手动添加的) // 自定义名称列表(手动添加的)
const allCustomNames = ref<string[]>([]) const allCustomNames = ref<string[]>([])
// 自定义名称列表 // 所有已知名称集合
const allKnownNames = computed(() => {
const set = new Set<string>()
upstreamModels.value.forEach(m => set.add(m.id))
return set
})
// 上游模型名称列表(去重后)
const upstreamModelNames = computed(() => {
const names = new Set<string>()
upstreamModels.value.forEach(m => {
names.add(m.id)
})
return Array.from(names).sort()
})
// 自定义名称列表(排除上游模型中已有的)
const customNames = computed(() => { const customNames = computed(() => {
return allCustomNames.value const upstreamSet = new Set(upstreamModelNames.value)
return allCustomNames.value.filter(name => !upstreamSet.has(name))
}) })
// 排序后的自定义名称 // 排序后的自定义名称
@@ -261,12 +371,26 @@ const canAddAsCustom = computed(() => {
if (!search) return false if (!search) return false
if (selectedNames.value.includes(search)) return false if (selectedNames.value.includes(search)) return false
if (allCustomNames.value.includes(search)) return false if (allCustomNames.value.includes(search)) return false
if (allKnownNames.value.has(search)) return false
return true return true
}) })
// 过滤后的上游模型
const filteredUpstreamModels = computed(() => {
if (!searchQuery.value.trim()) return upstreamModelNames.value
const query = searchQuery.value.toLowerCase()
return upstreamModelNames.value.filter(name => name.toLowerCase().includes(query))
})
// 空状态判断 // 空状态判断
const showEmptyState = computed(() => { const showEmptyState = computed(() => {
return customNames.value.length === 0 return filteredUpstreamModels.value.length === 0 && customNames.value.length === 0
})
// 上游模型是否全选
const isAllUpstreamModelsSelected = computed(() => {
if (filteredUpstreamModels.value.length === 0) return false
return filteredUpstreamModels.value.every(name => selectedNames.value.includes(name))
}) })
// 切换名称选中状态 // 切换名称选中状态
@@ -284,17 +408,71 @@ function addCustomName() {
const name = searchQuery.value.trim() const name = searchQuery.value.trim()
if (name && !selectedNames.value.includes(name)) { if (name && !selectedNames.value.includes(name)) {
selectedNames.value.push(name) selectedNames.value.push(name)
if (!allCustomNames.value.includes(name)) { if (!allKnownNames.value.has(name) && !allCustomNames.value.includes(name)) {
allCustomNames.value.push(name) allCustomNames.value.push(name)
} }
searchQuery.value = '' searchQuery.value = ''
} }
} }
// 全选/取消全选上游模型
function toggleAllUpstreamModels() {
const allNames = filteredUpstreamModels.value
if (isAllUpstreamModelsSelected.value) {
selectedNames.value = selectedNames.value.filter(name => !allNames.includes(name))
} else {
allNames.forEach(name => {
if (!selectedNames.value.includes(name)) {
selectedNames.value.push(name)
}
})
}
}
// 切换折叠状态
function toggleGroupCollapse(group: string) {
if (collapsedGroups.value.has(group)) {
collapsedGroups.value.delete(group)
} else {
collapsedGroups.value.add(group)
}
collapsedGroups.value = new Set(collapsedGroups.value)
}
// 从提供商获取模型(使用缓存)
async function fetchUpstreamModels() {
if (!props.providerId) return
try {
loadingModels.value = true
fetchingUpstreamModels.value = true
const result = await fetchCachedModels(props.providerId)
if (result.models.length > 0) {
upstreamModels.value = result.models
upstreamModelsLoaded.value = true
// 获取上游模型后,将不在上游列表中的已选名称添加到自定义列表
const upstreamIds = new Set(result.models.map(m => m.id))
const customFromSelected = selectedNames.value.filter(name => !upstreamIds.has(name))
const mergedCustom = new Set([...allCustomNames.value, ...customFromSelected])
allCustomNames.value = Array.from(mergedCustom).filter(name => !upstreamIds.has(name))
}
if (result.error) {
showError(result.error, '获取上游模型失败')
}
} catch (err: unknown) {
showError(parseApiError(err, '获取上游模型列表失败'), '错误')
} finally {
loadingModels.value = false
fetchingUpstreamModels.value = false
}
}
// 监听打开状态 // 监听打开状态
watch(() => props.open, async (isOpen) => { watch(() => props.open, async (isOpen) => {
if (isOpen) { if (isOpen) {
initForm() initForm()
if (props.hasAutoFetchKey) {
await fetchUpstreamModels()
}
} }
}) })
@@ -315,6 +493,9 @@ function initForm() {
allCustomNames.value = [] allCustomNames.value = []
} }
searchQuery.value = '' searchQuery.value = ''
upstreamModels.value = []
upstreamModelsLoaded.value = false
collapsedGroups.value = new Set()
} }
// 处理模型选择变更 // 处理模型选择变更

View File

@@ -541,7 +541,7 @@ async function loadAvailableGlobalModels() {
// 获取当前 provider 已添加的模型的 global_model_id 列表 // 获取当前 provider 已添加的模型的 global_model_id 列表
const existingGlobalModelIds = new Set( const existingGlobalModelIds = new Set(
existingModels.map((m: Model) => m.global_model_id).filter(Boolean) existingModels.map((m: Model) => m.global_model_id)
) )
// 过滤掉已添加的模型 // 过滤掉已添加的模型

View File

@@ -325,6 +325,7 @@
:models="models" :models="models"
:editing-group="editingGroup" :editing-group="editingGroup"
:preselected-model-id="preselectedModelId" :preselected-model-id="preselectedModelId"
:has-auto-fetch-key="hasAutoFetchKey"
@saved="onDialogSaved" @saved="onDialogSaved"
/> />
@@ -420,6 +421,11 @@ const aliasMappingPreview = computed(() => props.mappingPreview ?? null)
const providerEndpoints = computed(() => props.endpoints ?? []) const providerEndpoints = computed(() => props.endpoints ?? [])
const providerKeysState = computed(() => props.providerKeys ?? []) const providerKeysState = computed(() => props.providerKeys ?? [])
// 是否有 key 配置了自动获取上游模型
const hasAutoFetchKey = computed(() => {
return providerKeysState.value.some(k => k.auto_fetch_models)
})
// 展开状态 // 展开状态
const expandedItems = ref<Set<string>>(new Set()) const expandedItems = ref<Set<string>>(new Set())
@@ -641,18 +647,17 @@ async function onDialogSaved() {
emit('refresh') emit('refresh')
} }
// 获取可用的 API 格式(有活跃端点,去重) // 获取可用的 API 格式(有端点,去重;测试只关注 Key 是否支持,不依赖端点启用状态
const availableApiFormats = computed(() => { const availableApiFormats = computed(() => {
const formats = new Set( const formats = new Set(
providerEndpoints.value providerEndpoints.value
.filter(ep => ep.is_active)
.map(ep => ep.api_format) .map(ep => ep.api_format)
) )
return [...formats] return [...formats]
}) })
// 获取映射项支持的 API 格式 // 获取映射项支持的 API 格式
// 逻辑:找到支持该映射格式的所有活跃 Key获取这些 Key 支持的所有格式,与活跃端点格式取交集 // 逻辑:找到支持该映射格式的所有活跃 Key获取这些 Key 支持的所有格式,与端点格式取交集
function getItemAvailableFormats(item: CombinedMapping): string[] { function getItemAvailableFormats(item: CombinedMapping): string[] {
// 精确映射:基于 group.apiFormats 筛选 // 精确映射:基于 group.apiFormats 筛选
if (item.type === 'exact' && item.group?.apiFormats && item.group.apiFormats.length > 0) { if (item.type === 'exact' && item.group?.apiFormats && item.group.apiFormats.length > 0) {
@@ -677,7 +682,7 @@ function getItemAvailableFormats(item: CombinedMapping): string[] {
} }
} }
// 与活跃端点格式取交集 // 与端点格式取交集
return availableApiFormats.value.filter(fmt => keyFormats.has(fmt)) return availableApiFormats.value.filter(fmt => keyFormats.has(fmt))
} }

View File

@@ -281,6 +281,14 @@
<span class="text-xs font-mono">${{ (detail.cache_read_cost || 0).toFixed(6) }}</span> <span class="text-xs font-mono">${{ (detail.cache_read_cost || 0).toFixed(6) }}</span>
</div> </div>
</div> </div>
<!-- 缓存创建 5m/1h 细分 -->
<div
v-if="(detail.cache_creation_input_tokens_5m || 0) > 0 || (detail.cache_creation_input_tokens_1h || 0) > 0"
class="flex items-center pl-[56px]"
>
<span class="text-xs text-muted-foreground/50">5min: {{ detail.cache_creation_input_tokens_5m || 0 }}</span>
<span class="text-xs text-muted-foreground/50 ml-4">1h: {{ detail.cache_creation_input_tokens_1h || 0 }}</span>
</div>
</template> </template>
</div> </div>
</div> </div>

View File

@@ -502,6 +502,7 @@ const navigation = computed(() => {
title: '账户', title: '账户',
items: [ items: [
{ name: '使用统计', href: '/dashboard/usage', icon: BarChart3 }, { name: '使用统计', href: '/dashboard/usage', icon: BarChart3 },
{ name: '异步任务', href: '/dashboard/async-tasks', icon: Zap },
] ]
} }
] ]

View File

@@ -126,6 +126,11 @@ const routes: RouteRecordRaw[] = [
path: 'models', path: 'models',
name: 'ModelCatalog', name: 'ModelCatalog',
component: () => importWithRetry(() => import('@/views/user/ModelCatalog.vue')) component: () => importWithRetry(() => import('@/views/user/ModelCatalog.vue'))
},
{
path: 'async-tasks',
name: 'UserAsyncTasks',
component: () => importWithRetry(() => import('@/views/admin/AsyncTasks.vue'))
} }
] ]
}, },

View File

@@ -133,6 +133,16 @@ export function formatCost(cost: number | null | undefined): string {
return `$${cost.toFixed(4)}` return `$${cost.toFixed(4)}`
} }
// Usage count formatting (compact display for large numbers)
export function formatUsageCount(count: number): string {
if (count >= 1000000) {
return `${(count / 1000000).toFixed(1)}M`
} else if (count >= 1000) {
return `${(count / 1000).toFixed(1)}K`
}
return count.toString()
}
// Format remaining time from unix timestamp // Format remaining time from unix timestamp
export function formatRemainingTime(expireAt: number | undefined, currentTime: number): string { export function formatRemainingTime(expireAt: number | undefined, currentTime: number): string {
if (!expireAt) return '未知' if (!expireAt) return '未知'

View File

@@ -176,7 +176,7 @@
任务 任务
</TableHead> </TableHead>
<TableHead class="w-[15%]"> <TableHead class="w-[15%]">
用户/Provider {{ isAdmin ? '用户/Provider' : 'Provider' }}
</TableHead> </TableHead>
<TableHead class="w-[12%]"> <TableHead class="w-[12%]">
状态 状态
@@ -220,7 +220,10 @@
<!-- 用户/Provider --> <!-- 用户/Provider -->
<TableCell> <TableCell>
<div class="space-y-0.5 text-sm"> <div class="space-y-0.5 text-sm">
<div class="flex items-center gap-1.5"> <div
v-if="isAdmin"
class="flex items-center gap-1.5"
>
<User class="w-3 h-3 text-muted-foreground" /> <User class="w-3 h-3 text-muted-foreground" />
<span class="truncate max-w-[100px]">{{ task.username }}</span> <span class="truncate max-w-[100px]">{{ task.username }}</span>
</div> </div>
@@ -367,7 +370,10 @@
<!-- 信息网格 --> <!-- 信息网格 -->
<div class="grid grid-cols-2 gap-2 text-xs"> <div class="grid grid-cols-2 gap-2 text-xs">
<div class="flex items-center gap-1.5 text-muted-foreground"> <div
v-if="isAdmin"
class="flex items-center gap-1.5 text-muted-foreground"
>
<User class="w-3 h-3" /> <User class="w-3 h-3" />
<span class="truncate">{{ task.username }}</span> <span class="truncate">{{ task.username }}</span>
</div> </div>
@@ -492,8 +498,10 @@
</span> </span>
<span class="opacity-40">|</span> <span class="opacity-40">|</span>
<span>{{ formatDateFull(selectedTask.created_at) }}</span> <span>{{ formatDateFull(selectedTask.created_at) }}</span>
<span class="opacity-40">|</span> <template v-if="isAdmin">
<span>用户: {{ selectedTask.username }}</span> <span class="opacity-40">|</span>
<span>用户: {{ selectedTask.username }}</span>
</template>
<span class="opacity-40">|</span> <span class="opacity-40">|</span>
<span>Provider: {{ selectedTask.provider_name }}</span> <span>Provider: {{ selectedTask.provider_name }}</span>
</div> </div>
@@ -882,8 +890,11 @@ import {
ExternalLink, ExternalLink,
Copy, Copy,
} from 'lucide-vue-next' } from 'lucide-vue-next'
import { useAuthStore } from '@/stores/auth'
import { log } from '@/utils/logger' import { log } from '@/utils/logger'
const authStore = useAuthStore()
const isAdmin = computed(() => authStore.user?.role === 'admin')
const { toast } = useToast() const { toast } = useToast()
const { copyToClipboard } = useClipboard() const { copyToClipboard } = useClipboard()

View File

@@ -59,9 +59,6 @@
<TableHead class="w-[240px]"> <TableHead class="w-[240px]">
模型名称 模型名称
</TableHead> </TableHead>
<TableHead class="w-[140px]">
模型偏好
</TableHead>
<TableHead class="w-[160px] text-center"> <TableHead class="w-[160px] text-center">
价格 ($/M) 价格 ($/M)
</TableHead> </TableHead>
@@ -82,7 +79,7 @@
<TableBody> <TableBody>
<TableRow v-if="loading"> <TableRow v-if="loading">
<TableCell <TableCell
colspan="7" colspan="6"
class="text-center py-8" class="text-center py-8"
> >
<Loader2 class="w-6 h-6 animate-spin mx-auto" /> <Loader2 class="w-6 h-6 animate-spin mx-auto" />
@@ -90,7 +87,7 @@
</TableRow> </TableRow>
<TableRow v-else-if="filteredGlobalModels.length === 0"> <TableRow v-else-if="filteredGlobalModels.length === 0">
<TableCell <TableCell
colspan="7" colspan="6"
class="text-center py-8 text-muted-foreground" class="text-center py-8 text-muted-foreground"
> >
没有找到匹配的模型 没有找到匹配的模型
@@ -121,22 +118,6 @@
</div> </div>
</div> </div>
</TableCell> </TableCell>
<TableCell>
<div class="flex flex-wrap gap-0.5">
<template v-if="model.supported_capabilities?.length">
<span
v-for="capName in model.supported_capabilities"
:key="capName"
class="text-[11px] px-1 py-0.5 rounded bg-muted/60 text-muted-foreground"
:title="getCapabilityDisplayName(capName)"
>{{ getCapabilityShortName(capName) }}</span>
</template>
<span
v-else
class="text-muted-foreground text-xs"
>-</span>
</div>
</TableCell>
<TableCell class="text-center"> <TableCell class="text-center">
<div class="text-xs space-y-0.5"> <div class="text-xs space-y-0.5">
<!-- Token 计费 --> <!-- Token 计费 -->
@@ -297,19 +278,7 @@
</div> </div>
</div> </div>
<!-- 第二行模型偏好 --> <!-- 第二行统计信息 -->
<div
v-if="model.supported_capabilities?.length"
class="flex flex-wrap gap-0.5"
>
<span
v-for="capName in model.supported_capabilities"
:key="capName"
class="text-[11px] px-1 py-0.5 rounded bg-muted/60 text-muted-foreground"
>{{ getCapabilityShortName(capName) }}</span>
</div>
<!-- 第三行统计信息 -->
<div class="flex flex-wrap items-center gap-3 text-xs text-muted-foreground"> <div class="flex flex-wrap items-center gap-3 text-xs text-muted-foreground">
<span>提供商 {{ model.active_provider_count || 0 }}/{{ model.provider_count || 0 }}</span> <span>提供商 {{ model.active_provider_count || 0 }}/{{ model.provider_count || 0 }}</span>
<span>调用 {{ formatUsageCount(model.usage_count || 0) }}</span> <span>调用 {{ formatUsageCount(model.usage_count || 0) }}</span>
@@ -350,7 +319,6 @@
:model="selectedModel" :model="selectedModel"
:open="!!selectedModel" :open="!!selectedModel"
:has-blocking-dialog-open="hasBlockingDialogOpen" :has-blocking-dialog-open="hasBlockingDialogOpen"
:capabilities="capabilities"
@update:open="handleDrawerOpenChange" @update:open="handleDrawerOpenChange"
@edit-model="editModel" @edit-model="editModel"
@toggle-model-status="toggleModelStatus" @toggle-model-status="toggleModelStatus"
@@ -693,8 +661,8 @@ import {
type GlobalModelResponse, type GlobalModelResponse,
} from '@/api/global-models' } from '@/api/global-models'
import { log } from '@/utils/logger' import { log } from '@/utils/logger'
import { formatUsageCount } from '@/utils/format'
import { getProvidersSummary, type ProviderWithEndpointsSummary } from '@/api/endpoints/providers' import { getProvidersSummary, type ProviderWithEndpointsSummary } from '@/api/endpoints/providers'
import { getAllCapabilities, type CapabilityDefinition } from '@/api/endpoints'
interface ModelProviderDisplay { interface ModelProviderDisplay {
@@ -733,7 +701,6 @@ const editingModel = ref<GlobalModelResponse | null>(null)
// 数据 // 数据
const globalModels = ref<GlobalModelResponse[]>([]) const globalModels = ref<GlobalModelResponse[]>([])
const providers = ref<ProviderWithEndpointsSummary[]>([]) const providers = ref<ProviderWithEndpointsSummary[]>([])
const capabilities = ref<CapabilityDefinition[]>([])
// 模型目录分页 // 模型目录分页
const catalogCurrentPage = ref(1) const catalogCurrentPage = ref(1)
@@ -790,16 +757,6 @@ const editingProviderModel = computed<Model | null>(() => {
// 使用全局确认对话框 // 使用全局确认对话框
const { confirmDanger } = useConfirm() const { confirmDanger } = useConfirm()
// 格式化调用次数(大数字简化显示)
function formatUsageCount(count: number): string {
if (count >= 1000000) {
return `${(count / 1000000).toFixed(1) }M`
} else if (count >= 1000) {
return `${(count / 1000).toFixed(1) }K`
}
return count.toString()
}
// 从 GlobalModel 的 default_tiered_pricing 获取第一阶梯价格 // 从 GlobalModel 的 default_tiered_pricing 获取第一阶梯价格
function getFirstTierPrice(model: GlobalModelResponse, type: 'input' | 'output'): number | null { function getFirstTierPrice(model: GlobalModelResponse, type: 'input' | 'output'): number | null {
const tiered = model.default_tiered_pricing const tiered = model.default_tiered_pricing
@@ -1489,31 +1446,10 @@ async function loadProviders() {
} }
} }
async function loadCapabilities() {
try {
capabilities.value = await getAllCapabilities()
} catch (err) {
log.error('Failed to load capabilities:', err)
}
}
// 获取 capability 的显示名称
function getCapabilityDisplayName(capName: string): string {
const cap = capabilities.value.find(c => c.name === capName)
return cap?.display_name || capName
}
// 获取 capability 的短名称(用于表格展示)
function getCapabilityShortName(capName: string): string {
const cap = capabilities.value.find(c => c.name === capName)
return cap?.short_name || cap?.display_name || capName
}
onMounted(async () => { onMounted(async () => {
await Promise.all([ await Promise.all([
refreshData(), refreshData(),
loadProviders(), loadProviders(),
loadCapabilities(),
]) ])
}) })
</script> </script>

View File

@@ -652,7 +652,7 @@ async function loadOverview() {
overviewLoading.value = true overviewLoading.value = true
try { try {
const res = await getPoolOverview() const res = await getPoolOverview()
poolProviders.value = res.items poolProviders.value = res.items.filter(item => item.pool_enabled)
// Auto-select first provider if none selected // Auto-select first provider if none selected
if (!selectedProviderId.value && res.items.length > 0) { if (!selectedProviderId.value && res.items.length > 0) {
await selectProvider(res.items[0].provider_id) await selectProvider(res.items[0].provider_id)

View File

@@ -1,7 +1,7 @@
<template> <template>
<div class="space-y-6 px-4 sm:px-6 lg:px-0"> <div class="space-y-6 px-4 sm:px-6 lg:px-0">
<!-- 页面头部统计卡片 + 公告 --> <!-- 页面头部统计卡片 + 公告 -->
<div class="flex flex-col sm:flex-row gap-6 sm:items-start"> <div class="flex flex-col lg:flex-row gap-6 lg:items-start">
<!-- 左侧统计区域 --> <!-- 左侧统计区域 -->
<div <div
ref="statsPanelRef" ref="statsPanelRef"
@@ -215,12 +215,7 @@
Monthly Monthly
</Badge> </Badge>
</div> </div>
<div <div class="grid grid-cols-2 gap-2 sm:gap-3 xl:grid-cols-4">
class="grid gap-2 sm:gap-3"
:class="[
hasCacheData ? 'grid-cols-2 xl:grid-cols-4' : 'grid-cols-1 max-w-xs'
]"
>
<Card <Card
v-if="cacheStats" v-if="cacheStats"
class="relative p-3 sm:p-4 border-book-cloth/30" class="relative p-3 sm:p-4 border-book-cloth/30"
@@ -284,7 +279,7 @@
<!-- 右侧系统公告 --> <!-- 右侧系统公告 -->
<div <div
id="announcements-section" id="announcements-section"
class="w-full sm:w-[260px] md:w-[300px] lg:w-[320px] flex-shrink-0 flex flex-col min-h-0" class="w-full lg:w-[300px] xl:w-[320px] flex-shrink-0 flex flex-col min-h-0"
:style="announcementsContainerStyle" :style="announcementsContainerStyle"
> >
<div class="mb-3 flex items-center justify-between flex-shrink-0"> <div class="mb-3 flex items-center justify-between flex-shrink-0">
@@ -299,7 +294,7 @@
</Badge> </Badge>
</div> </div>
<Card class="overflow-hidden p-4 flex flex-col flex-1 min-h-0 h-full max-h-[280px] sm:max-h-none"> <Card class="overflow-hidden p-4 flex flex-col flex-1 min-h-0 h-full max-h-[280px] lg:max-h-none">
<div <div
v-if="loadingAnnouncements" v-if="loadingAnnouncements"
class="flex-1 flex items-center justify-center" class="flex-1 flex items-center justify-center"
@@ -856,7 +851,7 @@ const announcementsContainerStyle = computed(() => {
function checkScreenSize() { function checkScreenSize() {
if (typeof window !== 'undefined') { if (typeof window !== 'undefined') {
isLargeScreen.value = window.innerWidth >= 640 // sm breakpoint isLargeScreen.value = window.innerWidth >= 1024 // lg breakpoint
} }
} }

View File

@@ -40,12 +40,12 @@
<TableHead class="w-[140px] h-12 font-semibold"> <TableHead class="w-[140px] h-12 font-semibold">
模型名称 模型名称
</TableHead> </TableHead>
<TableHead class="w-[120px] h-12 font-semibold">
模型偏好
</TableHead>
<TableHead class="w-[140px] h-12 font-semibold text-center"> <TableHead class="w-[140px] h-12 font-semibold text-center">
价格 ($/M) 价格 ($/M)
</TableHead> </TableHead>
<TableHead class="w-[80px] h-12 font-semibold text-center">
调用次数
</TableHead>
<TableHead class="w-[70px] h-12 font-semibold text-center"> <TableHead class="w-[70px] h-12 font-semibold text-center">
状态 状态
</TableHead> </TableHead>
@@ -93,38 +93,6 @@
</div> </div>
</div> </div>
</TableCell> </TableCell>
<TableCell class="py-4">
<div class="flex gap-1.5 flex-wrap items-center">
<template v-if="getModelSupportedCapabilities(model).length > 0">
<button
v-for="cap in getModelSupportedCapabilitiesDetails(model)"
:key="cap.name"
class="inline-flex items-center gap-1 px-2 py-1 rounded-full text-xs font-medium transition-all"
:class="[
isCapabilityEnabled(model.name, cap.name)
? 'bg-primary text-primary-foreground'
: 'bg-transparent text-muted-foreground border border-dashed border-muted-foreground/50 hover:border-primary/50 hover:text-foreground'
]"
:title="cap.description"
@click.stop="toggleCapability(model.name, cap.name)"
>
<Check
v-if="isCapabilityEnabled(model.name, cap.name)"
class="w-3 h-3"
/>
<Plus
v-else
class="w-3 h-3"
/>
{{ cap.short_name || cap.display_name }}
</button>
</template>
<span
v-else
class="text-muted-foreground text-xs"
>-</span>
</div>
</TableCell>
<TableCell class="py-4 text-center"> <TableCell class="py-4 text-center">
<div class="text-xs space-y-0.5"> <div class="text-xs space-y-0.5">
<!-- Token 计费 --> <!-- Token 计费 -->
@@ -154,6 +122,9 @@
</div> </div>
</div> </div>
</TableCell> </TableCell>
<TableCell class="py-4 text-center">
<span class="text-sm font-mono">{{ formatUsageCount(model.usage_count || 0) }}</span>
</TableCell>
<TableCell class="py-4 text-center"> <TableCell class="py-4 text-center">
<Badge :variant="model.is_active ? 'success' : 'secondary'"> <Badge :variant="model.is_active ? 'success' : 'secondary'">
{{ model.is_active ? '可用' : '停用' }} {{ model.is_active ? '可用' : '停用' }}
@@ -194,41 +165,15 @@
</Badge> </Badge>
</div> </div>
<!-- 第二行价格 --> <!-- 第二行价格 + 调用次数 -->
<div <div class="flex items-center gap-3 text-xs text-muted-foreground">
v-if="getFirstTierPrice(model, 'input') || getFirstTierPrice(model, 'output')" <span
class="text-xs text-muted-foreground font-mono" v-if="getFirstTierPrice(model, 'input') || getFirstTierPrice(model, 'output')"
> class="font-mono"
In: ${{ getFirstTierPrice(model, 'input')?.toFixed(2) || '-' }} / Out: ${{ getFirstTierPrice(model, 'output')?.toFixed(2) || '-' }}
</div>
<!-- 第四行模型偏好按钮 -->
<div
v-if="getModelSupportedCapabilities(model).length > 0"
class="flex gap-1.5 flex-wrap"
@click.stop
>
<button
v-for="cap in getModelSupportedCapabilitiesDetails(model)"
:key="cap.name"
class="inline-flex items-center gap-1 px-2 py-1 rounded-full text-xs font-medium transition-all"
:class="[
isCapabilityEnabled(model.name, cap.name)
? 'bg-primary text-primary-foreground'
: 'bg-transparent text-muted-foreground border border-dashed border-muted-foreground/50'
]"
@click="toggleCapability(model.name, cap.name)"
> >
<Check In: ${{ getFirstTierPrice(model, 'input')?.toFixed(2) || '-' }} / Out: ${{ getFirstTierPrice(model, 'output')?.toFixed(2) || '-' }}
v-if="isCapabilityEnabled(model.name, cap.name)" </span>
class="w-3 h-3" <span class="font-mono">{{ formatUsageCount(model.usage_count || 0) }} </span>
/>
<Plus
v-else
class="w-3 h-3"
/>
{{ cap.short_name || cap.display_name }}
</button>
</div> </div>
</div> </div>
</div> </div>
@@ -250,10 +195,6 @@
<UserModelDetailDrawer <UserModelDetailDrawer
v-model:open="drawerOpen" v-model:open="drawerOpen"
:model="selectedModel" :model="selectedModel"
:capabilities="allCapabilities"
:user-configurable-capabilities="userConfigurableCapabilities"
:model-capability-settings="modelCapabilitySettings"
@toggle-capability="toggleCapability"
/> />
</div> </div>
</template> </template>
@@ -264,8 +205,6 @@ import {
Loader2, Loader2,
Search, Search,
Copy, Copy,
Check,
Plus,
} from 'lucide-vue-next' } from 'lucide-vue-next'
import { useToast } from '@/composables/useToast' import { useToast } from '@/composables/useToast'
import { useClipboard } from '@/composables/useClipboard' import { useClipboard } from '@/composables/useClipboard'
@@ -283,12 +222,8 @@ import {
RefreshButton, RefreshButton,
} from '@/components/ui' } from '@/components/ui'
import { type PublicGlobalModel } from '@/api/public-models' import { type PublicGlobalModel } from '@/api/public-models'
import { formatUsageCount } from '@/utils/format'
import { meApi } from '@/api/me' import { meApi } from '@/api/me'
import {
getUserConfigurableCapabilities,
getAllCapabilities,
type CapabilityDefinition
} from '@/api/endpoints'
import UserModelDetailDrawer from './components/UserModelDetailDrawer.vue' import UserModelDetailDrawer from './components/UserModelDetailDrawer.vue'
import { useRowClick } from '@/composables/useRowClick' import { useRowClick } from '@/composables/useRowClick'
import { log } from '@/utils/logger' import { log } from '@/utils/logger'
@@ -319,83 +254,6 @@ function openModelDetail(model: PublicGlobalModel, event: MouseEvent) {
const currentPage = ref(1) const currentPage = ref(1)
const pageSize = ref(20) const pageSize = ref(20)
// 能力筛选
const capabilityFilters = ref({
vision: false,
toolUse: false,
extendedThinking: false,
})
// 能力配置相关
const availableCapabilities = ref<CapabilityDefinition[]>([])
const allCapabilities = ref<CapabilityDefinition[]>([])
const userConfigurableCapabilities = computed(() =>
availableCapabilities.value.filter(cap => cap.config_mode === 'user_configurable')
)
const modelCapabilitySettings = ref<Record<string, Record<string, boolean>>>({})
const savingCapability = ref<string | null>(null) // 正在保存的能力标识 "modelName:capName"
// 获取模型支持的可配置能力名称列表(从 supported_capabilities 字段读取)
function getModelSupportedCapabilities(model: PublicGlobalModel): string[] {
if (!model.supported_capabilities) return []
// 只返回用户可配置的能力
return model.supported_capabilities.filter(capName =>
userConfigurableCapabilities.value.some(cap => cap.name === capName)
)
}
// 获取模型支持的可配置能力详情列表
function getModelSupportedCapabilitiesDetails(model: PublicGlobalModel): CapabilityDefinition[] {
const supportedNames = getModelSupportedCapabilities(model)
return userConfigurableCapabilities.value.filter(cap => supportedNames.includes(cap.name))
}
// 检查某个能力是否已启用
function isCapabilityEnabled(modelName: string, capName: string): boolean {
return modelCapabilitySettings.value[modelName]?.[capName] || false
}
// 切换能力配置
async function toggleCapability(modelName: string, capName: string) {
const capKey = `${modelName}:${capName}`
if (savingCapability.value === capKey) return // 防止重复点击
savingCapability.value = capKey
try {
const currentEnabled = isCapabilityEnabled(modelName, capName)
const newEnabled = !currentEnabled
// 更新本地状态
const newSettings = { ...modelCapabilitySettings.value }
if (!newSettings[modelName]) {
newSettings[modelName] = {}
}
if (newEnabled) {
newSettings[modelName][capName] = true
} else {
delete newSettings[modelName][capName]
// 如果该模型没有任何能力配置了,删除整个模型条目
if (Object.keys(newSettings[modelName]).length === 0) {
delete newSettings[modelName]
}
}
// 调用 API 保存
await meApi.updateModelCapabilitySettings({
model_capability_settings: Object.keys(newSettings).length > 0 ? newSettings : null
})
// 更新本地状态
modelCapabilitySettings.value = newSettings
} catch (err) {
log.error('保存能力配置失败:', err)
showError('保存失败,请重试')
} finally {
savingCapability.value = null
}
}
// 筛选后的模型列表 // 筛选后的模型列表
const filteredModels = computed(() => { const filteredModels = computed(() => {
let result = models.value let result = models.value
@@ -409,17 +267,6 @@ const filteredModels = computed(() => {
}) })
} }
// 能力筛选
if (capabilityFilters.value.vision) {
result = result.filter(m => m.config?.vision === true)
}
if (capabilityFilters.value.toolUse) {
result = result.filter(m => m.config?.function_calling === true)
}
if (capabilityFilters.value.extendedThinking) {
result = result.filter(m => m.config?.extended_thinking === true)
}
return result return result
}) })
@@ -430,10 +277,10 @@ const paginatedModels = computed(() => {
return filteredModels.value.slice(start, end) return filteredModels.value.slice(start, end)
}) })
// 搜索或筛选变化时重置到第一页 // 搜索变化时重置到第一页
watch([searchQuery, capabilityFilters], () => { watch(searchQuery, () => {
currentPage.value = 1 currentPage.value = 1
}, { deep: true }) })
async function loadModels() { async function loadModels() {
loading.value = true loading.value = true
@@ -449,30 +296,8 @@ async function loadModels() {
} }
} }
async function loadCapabilities() {
try {
const [userCaps, allCaps] = await Promise.all([
getUserConfigurableCapabilities(),
getAllCapabilities()
])
availableCapabilities.value = userCaps
allCapabilities.value = allCaps
} catch (err) {
log.error('Failed to load capabilities:', err)
}
}
async function loadModelCapabilitySettings() {
try {
const response = await meApi.getModelCapabilitySettings()
modelCapabilitySettings.value = response.model_capability_settings || {}
} catch (err) {
log.error('Failed to load model capability settings:', err)
}
}
async function refreshData() { async function refreshData() {
await Promise.all([loadModels(), loadCapabilities(), loadModelCapabilitySettings()]) await loadModels()
} }
// 从 PublicGlobalModel 的 default_tiered_pricing 获取第一阶梯价格 // 从 PublicGlobalModel 的 default_tiered_pricing 获取第一阶梯价格

View File

@@ -76,9 +76,6 @@
<TableHead class="min-w-[200px] h-12 font-semibold"> <TableHead class="min-w-[200px] h-12 font-semibold">
密钥名称 密钥名称
</TableHead> </TableHead>
<TableHead class="min-w-[80px] h-12 font-semibold">
能力
</TableHead>
<TableHead class="min-w-[160px] h-12 font-semibold"> <TableHead class="min-w-[160px] h-12 font-semibold">
密钥 密钥
</TableHead> </TableHead>
@@ -120,42 +117,6 @@
</div> </div>
</TableCell> </TableCell>
<!-- 能力 -->
<TableCell class="py-4">
<div class="flex gap-1.5 flex-wrap items-center">
<template v-if="userConfigurableCapabilities.length > 0">
<button
v-for="cap in userConfigurableCapabilities"
:key="cap.name"
class="inline-flex items-center gap-1 px-2 py-1 rounded-full text-xs font-medium transition-all"
:class="[
apiKey.is_locked ? 'opacity-50 cursor-not-allowed' : '',
isCapabilityEnabled(apiKey, cap.name)
? 'bg-primary text-primary-foreground'
: 'bg-transparent text-muted-foreground border border-dashed border-muted-foreground/50 hover:border-primary/50 hover:text-foreground'
]"
:title="apiKey.is_locked ? '已锁定' : getCapabilityTooltip(cap, isCapabilityEnabled(apiKey, cap.name))"
:disabled="apiKey.is_locked"
@click.stop="!apiKey.is_locked && toggleCapability(apiKey, cap.name)"
>
<Check
v-if="isCapabilityEnabled(apiKey, cap.name)"
class="w-3 h-3"
/>
<Plus
v-else
class="w-3 h-3"
/>
{{ cap.short_name || cap.display_name }}
</button>
</template>
<span
v-else
class="text-muted-foreground text-xs"
>-</span>
</div>
</TableCell>
<!-- 密钥显示 --> <!-- 密钥显示 -->
<TableCell class="py-4"> <TableCell class="py-4">
<div class="flex items-center gap-1.5"> <div class="flex items-center gap-1.5">
@@ -478,7 +439,6 @@
<script setup lang="ts"> <script setup lang="ts">
import { ref, onMounted } from 'vue' import { ref, onMounted } from 'vue'
import { meApi, type ApiKey } from '@/api/me' import { meApi, type ApiKey } from '@/api/me'
import { getAllCapabilities, type CapabilityDefinition } from '@/api/endpoints'
import Card from '@/components/ui/card.vue' import Card from '@/components/ui/card.vue'
import Button from '@/components/ui/button.vue' import Button from '@/components/ui/button.vue'
import Input from '@/components/ui/input.vue' import Input from '@/components/ui/input.vue'
@@ -495,7 +455,7 @@ import {
TableRow TableRow
} from '@/components/ui' } from '@/components/ui'
import RefreshButton from '@/components/ui/refresh-button.vue' import RefreshButton from '@/components/ui/refresh-button.vue'
import { Plus, Key, Copy, Trash2, Loader2, Activity, CheckCircle, Power, Check } from 'lucide-vue-next' import { Plus, Key, Copy, Trash2, Loader2, Activity, CheckCircle, Power } from 'lucide-vue-next'
import { useToast } from '@/composables/useToast' import { useToast } from '@/composables/useToast'
import { log } from '@/utils/logger' import { log } from '@/utils/logger'
import { parseApiError } from '@/utils/errorParser' import { parseApiError } from '@/utils/errorParser'
@@ -526,26 +486,10 @@ const newKeyName = ref('')
const newKeyValue = ref('') const newKeyValue = ref('')
const keyToDelete = ref<ApiKey | null>(null) const keyToDelete = ref<ApiKey | null>(null)
// 能力配置相关
const availableCapabilities = ref<CapabilityDefinition[]>([])
const userConfigurableCapabilities = computed(() =>
availableCapabilities.value.filter(cap => cap.config_mode === 'user_configurable')
)
const savingCapability = ref<string | null>(null) // 正在保存的能力标识 "keyId:capName"
onMounted(() => { onMounted(() => {
loadApiKeys() loadApiKeys()
loadCapabilities()
}) })
async function loadCapabilities() {
try {
availableCapabilities.value = await getAllCapabilities()
} catch (error) {
log.error('Failed to load capabilities:', error)
}
}
async function loadApiKeys() { async function loadApiKeys() {
loading.value = true loading.value = true
try { try {
@@ -625,50 +569,6 @@ async function toggleApiKey(apiKey: ApiKey) {
} }
} }
// 检查某个能力是否已启用
function isCapabilityEnabled(apiKey: ApiKey, capName: string): boolean {
return apiKey.force_capabilities?.[capName] || false
}
// 切换能力配置
async function toggleCapability(apiKey: ApiKey, capName: string) {
const capKey = `${apiKey.id}:${capName}`
if (savingCapability.value === capKey) return // 防止重复点击
savingCapability.value = capKey
try {
const currentEnabled = isCapabilityEnabled(apiKey, capName)
const newEnabled = !currentEnabled
// 构建新的能力配置
const newCapabilities: Record<string, boolean> = { ...(apiKey.force_capabilities || {}) }
if (newEnabled) {
newCapabilities[capName] = true
} else {
delete newCapabilities[capName]
}
const capabilitiesData = Object.keys(newCapabilities).length > 0 ? newCapabilities : null
// 调用 API 保存
await meApi.updateApiKeyCapabilities(apiKey.id, {
force_capabilities: capabilitiesData
})
// 更新本地数据
const index = apiKeys.value.findIndex(k => k.id === apiKey.id)
if (index !== -1) {
apiKeys.value[index].force_capabilities = capabilitiesData
}
} catch (err) {
log.error('保存能力配置失败:', err)
showError('保存失败,请重试')
} finally {
savingCapability.value = null
}
}
async function copyApiKey(apiKey: ApiKey) { async function copyApiKey(apiKey: ApiKey) {
try { try {
// 调用后端 API 获取完整密钥 // 调用后端 API 获取完整密钥
@@ -745,12 +645,4 @@ function formatRelativeTime(dateString: string): string {
return formatDate(dateString) return formatDate(dateString)
} }
// 获取能力按钮的提示文字
function getCapabilityTooltip(cap: CapabilityDefinition, isEnabled: boolean): string {
if (isEnabled) {
return `[已启用] 此密钥只能访问支持${cap.display_name}的模型`
}
return `${cap.description}`
}
</script> </script>

View File

@@ -150,51 +150,6 @@
</div> </div>
</div> </div>
<!-- 模型偏好 -->
<div
v-if="getModelUserConfigurableCapabilities().length > 0"
class="space-y-3"
>
<h4 class="font-semibold text-sm">
模型偏好
</h4>
<div class="space-y-2">
<div
v-for="cap in getModelUserConfigurableCapabilities()"
:key="cap.name"
class="flex items-center justify-between p-3 rounded-lg border"
>
<div class="flex-1 min-w-0">
<p class="text-sm font-medium">
{{ cap.display_name }}
</p>
<p
v-if="cap.description"
class="text-xs text-muted-foreground truncate"
>
{{ cap.description }}
</p>
</div>
<button
class="relative inline-flex h-5 w-9 shrink-0 cursor-pointer rounded-full border-2 border-transparent transition-colors duration-200 ease-in-out focus:outline-none focus-visible:ring-2 focus-visible:ring-ring focus-visible:ring-offset-2"
:class="[
isCapabilityEnabled(cap.name) ? 'bg-primary' : 'bg-muted'
]"
role="switch"
:aria-checked="isCapabilityEnabled(cap.name)"
@click="handleToggleCapability(cap.name)"
>
<span
class="pointer-events-none inline-block h-4 w-4 transform rounded-full bg-background shadow-lg ring-0 transition duration-200 ease-in-out"
:class="[
isCapabilityEnabled(cap.name) ? 'translate-x-4' : 'translate-x-0'
]"
/>
</button>
</div>
</div>
</div>
<!-- 定价信息 --> <!-- 定价信息 -->
<div class="space-y-3"> <div class="space-y-3">
<h4 class="font-semibold text-sm"> <h4 class="font-semibold text-sm">
@@ -365,13 +320,11 @@ import TableCell from '@/components/ui/table-cell.vue'
import type { PublicGlobalModel } from '@/api/public-models' import type { PublicGlobalModel } from '@/api/public-models'
import type { TieredPricingConfig, PricingTier } from '@/api/endpoints/types' import type { TieredPricingConfig, PricingTier } from '@/api/endpoints/types'
import type { CapabilityDefinition } from '@/api/endpoints'
const props = defineProps<Props>() const props = defineProps<Props>()
const emit = defineEmits<{ const emit = defineEmits<{
'update:open': [value: boolean] 'update:open': [value: boolean]
'toggleCapability': [modelName: string, capName: string]
}>() }>()
const { copyToClipboard } = useClipboard() const { copyToClipboard } = useClipboard()
@@ -379,29 +332,6 @@ const { copyToClipboard } = useClipboard()
interface Props { interface Props {
model: PublicGlobalModel | null model: PublicGlobalModel | null
open: boolean open: boolean
capabilities?: CapabilityDefinition[]
userConfigurableCapabilities?: CapabilityDefinition[]
modelCapabilitySettings?: Record<string, Record<string, boolean>>
}
// 获取模型支持的用户可配置能力
function getModelUserConfigurableCapabilities(): CapabilityDefinition[] {
if (!props.model?.supported_capabilities || !props.userConfigurableCapabilities) return []
return props.userConfigurableCapabilities.filter(cap =>
props.model?.supported_capabilities?.includes(cap.name)
)
}
// 检查能力是否已启用
function isCapabilityEnabled(capName: string): boolean {
if (!props.model) return false
return props.modelCapabilitySettings?.[props.model.name]?.[capName] || false
}
// 切换能力
function handleToggleCapability(capName: string) {
if (!props.model) return
emit('toggleCapability', props.model.name, capName)
} }
function handleClose() { function handleClose() {

View File

@@ -30,6 +30,7 @@ from src.services.model.fetch_scheduler import (
from src.services.model.upstream_fetcher import ( from src.services.model.upstream_fetcher import (
EndpointFetchConfig, EndpointFetchConfig,
UpstreamModelsFetchContext, UpstreamModelsFetchContext,
UpstreamModelsFetcherRegistry,
build_format_to_config, build_format_to_config,
fetch_models_for_key, fetch_models_for_key,
get_adapter_for_format, get_adapter_for_format,
@@ -235,7 +236,15 @@ async def query_available_models(
# 构建 api_format -> EndpointFetchConfig 映射(纯数据,不依赖 ORM session # 构建 api_format -> EndpointFetchConfig 映射(纯数据,不依赖 ORM session
format_to_endpoint = build_format_to_config(provider.endpoints) format_to_endpoint = build_format_to_config(provider.endpoints)
if not format_to_endpoint: # 检查是否有注册自定义 fetcher如预设模型有则不依赖活跃 endpoint
provider_type = str(getattr(provider, "provider_type", "") or "").lower()
# 延迟导入避免循环依赖(与 upstream_fetcher.fetch_models_for_key 保持一致)
from src.services.provider.envelope import ensure_providers_bootstrapped
ensure_providers_bootstrapped()
has_custom_fetcher = UpstreamModelsFetcherRegistry.get(provider_type) is not None
if not format_to_endpoint and not has_custom_fetcher:
raise HTTPException(status_code=400, detail="No active endpoints found for this provider") raise HTTPException(status_code=400, detail="No active endpoints found for this provider")
# 如果指定了 api_key_id只获取该 Key 的模型 # 如果指定了 api_key_id只获取该 Key 的模型
@@ -253,7 +262,6 @@ async def query_available_models(
raise HTTPException(status_code=400, detail="No active API Key found for this provider") raise HTTPException(status_code=400, detail="No active API Key found for this provider")
# Antigravity: 按 tier/可用性排序后逐个尝试,成功即停止 # Antigravity: 按 tier/可用性排序后逐个尝试,成功即停止
provider_type = str(getattr(provider, "provider_type", "") or "").lower()
if provider_type == ProviderType.ANTIGRAVITY: if provider_type == ProviderType.ANTIGRAVITY:
return await _fetch_models_antigravity_ordered( return await _fetch_models_antigravity_ordered(
provider=provider, provider=provider,
@@ -610,12 +618,12 @@ async def test_model(
raise HTTPException(status_code=404, detail="Provider not found") raise HTTPException(status_code=404, detail="Provider not found")
# 构建 api_format -> endpoint 映射 和 id -> endpoint 映射 # 构建 api_format -> endpoint 映射 和 id -> endpoint 映射
# 测试不依赖端点启用状态,禁用的端点也可以用于测试连通性
format_to_endpoint: dict[str, ProviderEndpoint] = {} format_to_endpoint: dict[str, ProviderEndpoint] = {}
id_to_endpoint: dict[str, ProviderEndpoint] = {} id_to_endpoint: dict[str, ProviderEndpoint] = {}
for ep in provider.endpoints: for ep in provider.endpoints:
if ep.is_active: format_to_endpoint[ep.api_format] = ep
format_to_endpoint[ep.api_format] = ep id_to_endpoint[ep.id] = ep
id_to_endpoint[ep.id] = ep
# 找到合适的端点和 API Key # 找到合适的端点和 API Key
endpoint = None endpoint = None
@@ -628,7 +636,7 @@ async def test_model(
if not endpoint: if not endpoint:
raise HTTPException( raise HTTPException(
status_code=404, status_code=404,
detail=f"No active endpoint found for API format: {request.api_format}", detail=f"No endpoint found for API format: {request.api_format}",
) )
if request.api_key_id: if request.api_key_id:
@@ -657,7 +665,7 @@ async def test_model(
# 使用指定的端点 # 使用指定的端点
endpoint = id_to_endpoint.get(request.endpoint_id) endpoint = id_to_endpoint.get(request.endpoint_id)
if not endpoint: if not endpoint:
raise HTTPException(status_code=404, detail="Endpoint not found or not active") raise HTTPException(status_code=404, detail="Endpoint not found")
if request.api_key_id: if request.api_key_id:
# 同时指定了 Key需要校验是否支持该端点格式 # 同时指定了 Key需要校验是否支持该端点格式

View File

@@ -715,6 +715,7 @@ class AdminImportFromUpstreamAdapter(AdminApiAdapter):
# 1. 检查是否已存在同名的 ProviderModel # 1. 检查是否已存在同名的 ProviderModel
existing = ( existing = (
db.query(Model) db.query(Model)
.options(joinedload(Model.global_model))
.filter( .filter(
Model.provider_id == self.provider_id, Model.provider_id == self.provider_id,
Model.provider_model_name == model_id, Model.provider_model_name == model_id,
@@ -727,10 +728,8 @@ class AdminImportFromUpstreamAdapter(AdminApiAdapter):
success.append( success.append(
ImportFromUpstreamSuccessItem( ImportFromUpstreamSuccessItem(
model_id=model_id, model_id=model_id,
global_model_id=existing.global_model_id or "", global_model_id=existing.global_model_id,
global_model_name=( global_model_name=existing.global_model.name,
existing.global_model.name if existing.global_model else ""
),
provider_model_id=existing.id, provider_model_id=existing.id,
created_global_model=False, created_global_model=False,
) )

View File

@@ -319,7 +319,6 @@ def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndp
.filter( .filter(
Model.provider_id == provider.id, Model.provider_id == provider.id,
Model.is_active == True, Model.is_active == True,
Model.global_model_id.isnot(None),
) )
.distinct() .distinct()
.all() .all()

View File

@@ -2556,16 +2556,18 @@ def _purge_stats_and_reset_counters(db: Session) -> None:
class AdminPurgeUsageAdapter(AdminApiAdapter): class AdminPurgeUsageAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override] async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
"""清空全部使用记录及相关统计数据""" """清空全部使用记录及相关统计数据"""
from src.models.database import RequestCandidate from src.models.database import RequestCandidate, UserModelUsageCount
db = context.db db = context.db
usage_count = db.query(Usage).count() usage_count = db.query(Usage).count()
candidates_count = db.query(RequestCandidate).count() candidates_count = db.query(RequestCandidate).count()
usage_counts_count = db.query(UserModelUsageCount).count()
# 清空使用记录 # 清空使用记录
db.query(RequestCandidate).delete() db.query(RequestCandidate).delete()
db.query(Usage).delete() db.query(Usage).delete()
db.query(UserModelUsageCount).delete()
_purge_stats_and_reset_counters(db) _purge_stats_and_reset_counters(db)
db.commit() db.commit()
@@ -2575,6 +2577,7 @@ class AdminPurgeUsageAdapter(AdminApiAdapter):
"deleted": { "deleted": {
"usage_records": usage_count, "usage_records": usage_count,
"request_candidates": candidates_count, "request_candidates": candidates_count,
"user_model_usage_counts": usage_counts_count,
}, },
} }

View File

@@ -1259,6 +1259,8 @@ class AdminUsageDetailAdapter(AdminApiAdapter):
}, },
"cache_creation_input_tokens": usage_record.cache_creation_input_tokens, "cache_creation_input_tokens": usage_record.cache_creation_input_tokens,
"cache_read_input_tokens": usage_record.cache_read_input_tokens, "cache_read_input_tokens": usage_record.cache_read_input_tokens,
"cache_creation_input_tokens_5m": usage_record.cache_creation_input_tokens_5m or 0,
"cache_creation_input_tokens_1h": usage_record.cache_creation_input_tokens_1h or 0,
"cache_creation_cost": getattr(usage_record, "cache_creation_cost_usd", 0.0), "cache_creation_cost": getattr(usage_record, "cache_creation_cost_usd", 0.0),
"cache_read_cost": getattr(usage_record, "cache_read_cost_usd", 0.0), "cache_read_cost": getattr(usage_record, "cache_read_cost_usd", 0.0),
"request_cost": getattr(usage_record, "request_cost_usd", 0.0), "request_cost": getattr(usage_record, "request_cost_usd", 0.0),

View File

@@ -184,6 +184,8 @@ class ChatSyncExecutor:
output_tokens = usage_info.get("output_tokens", 0) output_tokens = usage_info.get("output_tokens", 0)
cache_creation_tokens = usage_info.get("cache_creation_input_tokens", 0) cache_creation_tokens = usage_info.get("cache_creation_input_tokens", 0)
cached_tokens = usage_info.get("cache_read_input_tokens", 0) cached_tokens = usage_info.get("cache_read_input_tokens", 0)
cache_creation_tokens_5m = usage_info.get("cache_creation_input_tokens_5m", 0)
cache_creation_tokens_1h = usage_info.get("cache_creation_input_tokens_1h", 0)
# 非流式成功时,返回给客户端的是提供商响应头(透传) # 非流式成功时,返回给客户端的是提供商响应头(透传)
# JSONResponse 会自动设置 content-type但我们记录实际返回的完整头 # JSONResponse 会自动设置 content-type但我们记录实际返回的完整头
@@ -211,6 +213,8 @@ class ChatSyncExecutor:
provider_request_body=ctx.provider_request_body, provider_request_body=ctx.provider_request_body,
cache_creation_tokens=cache_creation_tokens, cache_creation_tokens=cache_creation_tokens,
cache_read_tokens=cached_tokens, cache_read_tokens=cached_tokens,
cache_creation_tokens_5m=cache_creation_tokens_5m,
cache_creation_tokens_1h=cache_creation_tokens_1h,
is_stream=False, is_stream=False,
provider_request_headers=ctx.provider_request_headers, provider_request_headers=ctx.provider_request_headers,
api_format=api_format, api_format=api_format,

View File

@@ -171,7 +171,7 @@ async def _calculate_and_record_usage(
provider = db.query(Provider).filter(Provider.id == provider_api_key.provider_id).first() provider = db.query(Provider).filter(Provider.id == provider_api_key.provider_id).first()
if provider: if provider:
for ep in provider.endpoints: for ep in provider.endpoints:
if ep.api_format == api_format and ep.is_active: if ep.api_format == api_format:
provider_endpoint = ep provider_endpoint = ep
break break

View File

@@ -86,6 +86,8 @@ class StreamContext:
output_tokens: int = 0 output_tokens: int = 0
cached_tokens: int = 0 cached_tokens: int = 0
cache_creation_tokens: int = 0 cache_creation_tokens: int = 0
cache_creation_tokens_5m: int = 0 # 5min TTL 缓存创建
cache_creation_tokens_1h: int = 0 # 1h TTL 缓存创建
# 响应内容 # 响应内容
_collected_text_parts: list[str] = field(default_factory=list, repr=False) _collected_text_parts: list[str] = field(default_factory=list, repr=False)
@@ -159,6 +161,8 @@ class StreamContext:
self.output_tokens = 0 self.output_tokens = 0
self.cached_tokens = 0 self.cached_tokens = 0
self.cache_creation_tokens = 0 self.cache_creation_tokens = 0
self.cache_creation_tokens_5m = 0
self.cache_creation_tokens_1h = 0
self.error_message = None self.error_message = None
self.upstream_response = None self.upstream_response = None
self.status_code = 200 self.status_code = 200

View File

@@ -228,6 +228,8 @@ class StreamTelemetryRecorder:
provider_request_body=ctx.provider_request_body, provider_request_body=ctx.provider_request_body,
cache_creation_tokens=ctx.cache_creation_tokens, cache_creation_tokens=ctx.cache_creation_tokens,
cache_read_tokens=ctx.cached_tokens, cache_read_tokens=ctx.cached_tokens,
cache_creation_tokens_5m=ctx.cache_creation_tokens_5m,
cache_creation_tokens_1h=ctx.cache_creation_tokens_1h,
is_stream=True, is_stream=True,
provider_request_headers=ctx.provider_request_headers, provider_request_headers=ctx.provider_request_headers,
api_format=ctx.api_format, api_format=ctx.api_format,
@@ -285,6 +287,8 @@ class StreamTelemetryRecorder:
output_tokens=ctx.output_tokens, output_tokens=ctx.output_tokens,
cache_creation_tokens=ctx.cache_creation_tokens, cache_creation_tokens=ctx.cache_creation_tokens,
cache_read_tokens=ctx.cached_tokens, cache_read_tokens=ctx.cached_tokens,
cache_creation_tokens_5m=ctx.cache_creation_tokens_5m,
cache_creation_tokens_1h=ctx.cache_creation_tokens_1h,
response_body=response_body, response_body=response_body,
client_response_body=client_response_body, client_response_body=client_response_body,
response_headers=ctx.response_headers, response_headers=ctx.response_headers,
@@ -342,6 +346,8 @@ class StreamTelemetryRecorder:
output_tokens=ctx.output_tokens, output_tokens=ctx.output_tokens,
cache_creation_tokens=ctx.cache_creation_tokens, cache_creation_tokens=ctx.cache_creation_tokens,
cache_read_tokens=ctx.cached_tokens, cache_read_tokens=ctx.cached_tokens,
cache_creation_tokens_5m=ctx.cache_creation_tokens_5m,
cache_creation_tokens_1h=ctx.cache_creation_tokens_1h,
response_body=response_body, response_body=response_body,
client_response_body=client_response_body, client_response_body=client_response_body,
response_headers=ctx.response_headers, response_headers=ctx.response_headers,

View File

@@ -94,6 +94,34 @@ def extract_cache_creation_tokens(usage: dict[str, Any]) -> int:
return old_format return old_format
def extract_cache_creation_tokens_detail(usage: dict[str, Any]) -> tuple[int, int, int]:
"""
提取缓存创建 tokens 细分(区分 5m 和 1h
返回 (total, tokens_5m, tokens_1h) 三元组。
当无法区分时tokens_5m 和 tokens_1h 均为 0total 为合计值。
"""
# 1. 嵌套格式
cache_creation = usage.get("cache_creation")
if isinstance(cache_creation, dict) and (
"ephemeral_5m_input_tokens" in cache_creation
or "ephemeral_1h_input_tokens" in cache_creation
):
t5m = int(cache_creation.get("ephemeral_5m_input_tokens", 0))
t1h = int(cache_creation.get("ephemeral_1h_input_tokens", 0))
return t5m + t1h, t5m, t1h
# 2. 扁平新格式
if "claude_cache_creation_5_m_tokens" in usage or "claude_cache_creation_1_h_tokens" in usage:
t5m = int(usage.get("claude_cache_creation_5_m_tokens", 0))
t1h = int(usage.get("claude_cache_creation_1_h_tokens", 0))
return t5m + t1h, t5m, t1h
# 3. 旧格式:无法区分
old = int(usage.get("cache_creation_input_tokens", 0))
return old, 0, 0
def build_sse_headers(extra_headers: dict[str, str] | None = None) -> dict[str, str]: def build_sse_headers(extra_headers: dict[str, str] | None = None) -> dict[str, str]:
""" """
构建 SSEtext/event-stream推荐响应头用于减少代理缓冲带来的卡顿/成段输出。 构建 SSEtext/event-stream推荐响应头用于减少代理缓冲带来的卡顿/成段输出。

View File

@@ -31,14 +31,15 @@ class ClaudeCapabilityDetector:
request_body: dict[str, Any] | None = None, request_body: dict[str, Any] | None = None,
) -> dict[str, bool]: ) -> dict[str, bool]:
""" """
从 Claude 请求头检测能力需求 从 Claude 请求头和请求体检测能力需求
检测规则: 检测规则:
- anthropic-beta: context-1m-xxx -> context_1m: True - anthropic-beta: context-1m-xxx -> context_1m: True
- 请求体中 cache_control.ttl = "1h" -> cache_1h: True
Args: Args:
headers: 请求头字典 headers: 请求头字典
request_body: 请求体(Claude 不使用,保留用于接口统一 request_body: 请求体(用于检测 cache_control.ttl
""" """
requirements: dict[str, bool] = {} requirements: dict[str, bool] = {}
@@ -47,9 +48,59 @@ class ClaudeCapabilityDetector:
if beta_header and "context-1m" in beta_header.lower(): if beta_header and "context-1m" in beta_header.lower():
requirements["context_1m"] = True requirements["context_1m"] = True
# 从请求体检测 cache_1h
if request_body and _detect_cache_1h_in_body(request_body):
requirements["cache_1h"] = True
return requirements return requirements
def _has_cache_1h_ttl(block: dict[str, Any]) -> bool:
"""检查单个内容块是否包含 cache_control.ttl = '1h'"""
cache_control = block.get("cache_control")
if isinstance(cache_control, dict):
return cache_control.get("ttl") == "1h"
return False
def _detect_cache_1h_in_body(body: dict[str, Any]) -> bool:
"""
扫描 Claude 请求体,检测是否包含 cache_control.ttl = "1h"
检查位置:
- system[].cache_control.ttl
- messages[].content[].cache_control.ttl
- tools[].cache_control.ttl
"""
# 检查 system数组格式
system = body.get("system")
if isinstance(system, list):
for block in system:
if isinstance(block, dict) and _has_cache_1h_ttl(block):
return True
# 检查 messages
messages = body.get("messages")
if isinstance(messages, list):
for msg in messages:
if not isinstance(msg, dict):
continue
content = msg.get("content")
if isinstance(content, list):
for block in content:
if isinstance(block, dict) and _has_cache_1h_ttl(block):
return True
# 检查 tools
tools = body.get("tools")
if isinstance(tools, list):
for tool in tools:
if isinstance(tool, dict) and _has_cache_1h_ttl(tool):
return True
return False
@register_adapter @register_adapter
class ClaudeChatAdapter(ChatAdapterBase): class ClaudeChatAdapter(ChatAdapterBase):
""" """

View File

@@ -8,7 +8,7 @@ Claude Chat Handler - 基于通用 Chat Handler 基类的简化实现
from typing import Any from typing import Any
from src.api.handlers.base.chat_handler_base import ChatHandlerBase from src.api.handlers.base.chat_handler_base import ChatHandlerBase
from src.api.handlers.base.utils import extract_cache_creation_tokens from src.api.handlers.base.utils import extract_cache_creation_tokens_detail
from src.core.api_format import ApiFamily, EndpointKind from src.core.api_format import ApiFamily, EndpointKind
@@ -103,12 +103,15 @@ class ClaudeChatHandler(ChatHandlerBase):
- 新格式claude_cache_creation_5_m_tokens / claude_cache_creation_1_h_tokens - 新格式claude_cache_creation_5_m_tokens / claude_cache_creation_1_h_tokens
""" """
usage = response.get("usage", {}) usage = response.get("usage", {})
total, t5m, t1h = extract_cache_creation_tokens_detail(usage)
return { return {
"input_tokens": usage.get("input_tokens", 0), "input_tokens": usage.get("input_tokens", 0),
"output_tokens": usage.get("output_tokens", 0), "output_tokens": usage.get("output_tokens", 0),
"cache_creation_input_tokens": extract_cache_creation_tokens(usage), "cache_creation_input_tokens": total,
"cache_read_input_tokens": usage.get("cache_read_input_tokens", 0), "cache_read_input_tokens": usage.get("cache_read_input_tokens", 0),
"cache_creation_input_tokens_5m": t5m,
"cache_creation_input_tokens_1h": t1h,
} }
def _normalize_response(self, response: dict[str, Any]) -> dict[str, Any]: def _normalize_response(self, response: dict[str, Any]) -> dict[str, Any]:

View File

@@ -46,7 +46,7 @@ class ClaudeCliAdapter(CliAdapterBase):
request_body: dict[str, Any] | None = None, request_body: dict[str, Any] | None = None,
) -> dict[str, bool]: ) -> dict[str, bool]:
"""检测 Claude CLI 请求中隐含的能力需求""" """检测 Claude CLI 请求中隐含的能力需求"""
return ClaudeCapabilityDetector.detect_from_headers(headers) return ClaudeCapabilityDetector.detect_from_headers(headers, request_body)
# ========================================================================= # =========================================================================
# Claude CLI 特定的计费逻辑 # Claude CLI 特定的计费逻辑

View File

@@ -10,7 +10,7 @@ from src.api.handlers.base.cli_handler_base import (
CliMessageHandlerBase, CliMessageHandlerBase,
StreamContext, StreamContext,
) )
from src.api.handlers.base.utils import extract_cache_creation_tokens from src.api.handlers.base.utils import extract_cache_creation_tokens_detail
from src.core.api_format import ApiFamily, EndpointKind from src.core.api_format import ApiFamily, EndpointKind
@@ -114,9 +114,11 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
if cache_read: if cache_read:
ctx.cached_tokens = cache_read ctx.cached_tokens = cache_read
cache_creation = extract_cache_creation_tokens(usage) total, t5m, t1h = extract_cache_creation_tokens_detail(usage)
if cache_creation: if total:
ctx.cache_creation_tokens = cache_creation ctx.cache_creation_tokens = total
ctx.cache_creation_tokens_5m = t5m
ctx.cache_creation_tokens_1h = t1h
# 处理文本增量 # 处理文本增量
elif event_type == "content_block_delta": elif event_type == "content_block_delta":
@@ -140,9 +142,11 @@ class ClaudeCliMessageHandler(CliMessageHandlerBase):
ctx.cached_tokens = usage["cache_read_input_tokens"] ctx.cached_tokens = usage["cache_read_input_tokens"]
# 更新缓存创建 tokens # 更新缓存创建 tokens
cache_creation = extract_cache_creation_tokens(usage) total, t5m, t1h = extract_cache_creation_tokens_detail(usage)
if cache_creation > 0: if total > 0:
ctx.cache_creation_tokens = cache_creation ctx.cache_creation_tokens = total
ctx.cache_creation_tokens_5m = t5m
ctx.cache_creation_tokens_1h = t1h
# 检查是否结束 # 检查是否结束
delta = data.get("delta", {}) delta = data.get("delta", {})

View File

@@ -19,9 +19,30 @@ from src.core.api_format.enums import AuthMethod
from src.core.api_format.headers import BROWSER_FINGERPRINT_HEADERS from src.core.api_format.headers import BROWSER_FINGERPRINT_HEADERS
from src.core.logger import logger from src.core.logger import logger
from src.models.gemini import GeminiRequest from src.models.gemini import GeminiRequest
from src.services.gemini_files_mapping import extract_file_names_from_request
from src.services.provider.transport import redact_url_for_log from src.services.provider.transport import redact_url_for_log
class GeminiCapabilityDetector:
"""Gemini API 能力检测器"""
@staticmethod
def detect_from_request(
headers: dict[str, str], # noqa: ARG004 - 预留
request_body: dict[str, Any] | None = None,
) -> dict[str, bool]:
"""
从请求体检测 Gemini 能力需求
检测规则:
- fileData.fileUri -> gemini_files: True
"""
requirements: dict[str, bool] = {}
if request_body and extract_file_names_from_request(request_body):
requirements["gemini_files"] = True
return requirements
@register_adapter @register_adapter
class GeminiChatAdapter(ChatAdapterBase): class GeminiChatAdapter(ChatAdapterBase):
""" """
@@ -62,11 +83,11 @@ class GeminiChatAdapter(ChatAdapterBase):
def detect_capability_requirements( def detect_capability_requirements(
self, self,
headers: dict[str, str], # noqa: ARG002 - 预留 headers: dict[str, str],
request_body: dict[str, Any] | None = None, # noqa: ARG002 - 预留 request_body: dict[str, Any] | None = None,
) -> dict[str, bool]: ) -> dict[str, bool]:
"""Gemini API 无特殊能力要求""" """从请求体检测 Gemini 能力需求fileData.fileUri -> gemini_files"""
return {} return GeminiCapabilityDetector.detect_from_request(headers, request_body)
def _merge_path_params( def _merge_path_params(
self, original_request_body: dict[str, Any], path_params: dict[str, Any] # noqa: ARG002 self, original_request_body: dict[str, Any], path_params: dict[str, Any] # noqa: ARG002

View File

@@ -13,7 +13,7 @@ from fastapi import Request
from src.api.handlers.base.cli_adapter_base import CliAdapterBase, register_cli_adapter from src.api.handlers.base.cli_adapter_base import CliAdapterBase, register_cli_adapter
from src.api.handlers.base.cli_handler_base import CliMessageHandlerBase from src.api.handlers.base.cli_handler_base import CliMessageHandlerBase
from src.api.handlers.gemini.adapter import GeminiChatAdapter from src.api.handlers.gemini.adapter import GeminiCapabilityDetector, GeminiChatAdapter
from src.config.settings import config from src.config.settings import config
from src.core.api_format import ApiFamily, get_auth_handler from src.core.api_format import ApiFamily, get_auth_handler
from src.core.api_format.enums import AuthMethod from src.core.api_format.enums import AuthMethod
@@ -53,6 +53,14 @@ class GeminiCliAdapter(CliAdapterBase):
handler = get_auth_handler(AuthMethod.GOOG_API_KEY) handler = get_auth_handler(AuthMethod.GOOG_API_KEY)
return handler.extract_credentials(request) return handler.extract_credentials(request)
def detect_capability_requirements(
self,
headers: dict[str, str],
request_body: dict[str, Any] | None = None,
) -> dict[str, bool]:
"""从请求体检测 Gemini 能力需求fileData.fileUri -> gemini_files"""
return GeminiCapabilityDetector.detect_from_request(headers, request_body)
def _merge_path_params( def _merge_path_params(
self, original_request_body: dict[str, Any], path_params: dict[str, Any] # noqa: ARG002 self, original_request_body: dict[str, Any], path_params: dict[str, Any] # noqa: ARG002
) -> dict[str, Any]: ) -> dict[str, Any]:

View File

@@ -312,18 +312,13 @@ class PublicProvidersAdapter(PublicApiAdapter):
providers = query.offset(self.skip).limit(self.limit).all() providers = query.offset(self.skip).limit(self.limit).all()
result = [] result = []
for provider in providers: for provider in providers:
models_count = ( models_count = db.query(Model).filter(Model.provider_id == provider.id).count()
db.query(Model)
.filter(Model.provider_id == provider.id, Model.global_model_id.isnot(None))
.count()
)
active_models_count = ( active_models_count = (
db.query(Model) db.query(Model)
.filter( .filter(
and_( and_(
Model.provider_id == provider.id, Model.provider_id == provider.id,
Model.is_active.is_(True), Model.is_active.is_(True),
Model.global_model_id.isnot(None),
) )
) )
.count() .count()
@@ -367,7 +362,6 @@ class PublicModelsAdapter(PublicApiAdapter):
and_( and_(
Model.is_active.is_(True), Model.is_active.is_(True),
Provider.is_active.is_(True), Provider.is_active.is_(True),
Model.global_model_id.isnot(None),
) )
) )
) )
@@ -424,7 +418,6 @@ class PublicStatsAdapter(PublicApiAdapter):
and_( and_(
Model.is_active.is_(True), Model.is_active.is_(True),
Provider.is_active.is_(True), Provider.is_active.is_(True),
Model.global_model_id.isnot(None),
) )
) )
.count() .count()
@@ -462,7 +455,6 @@ class PublicSearchModelsAdapter(PublicApiAdapter):
and_( and_(
Model.is_active.is_(True), Model.is_active.is_(True),
Provider.is_active.is_(True), Provider.is_active.is_(True),
Model.global_model_id.isnot(None),
) )
) )
) )

View File

@@ -32,7 +32,15 @@ from src.models.api import (
UpdatePreferencesRequest, UpdatePreferencesRequest,
UpdateProfileRequest, UpdateProfileRequest,
) )
from src.models.database import ApiKey, GlobalModel, Model, Provider, Usage, User from src.models.database import (
ApiKey,
GlobalModel,
Model,
Provider,
Usage,
User,
UserModelUsageCount,
)
from src.services.system.time_range import TimeRangeParams from src.services.system.time_range import TimeRangeParams
from src.services.usage.service import UsageService from src.services.usage.service import UsageService
from src.services.user.apikey import ApiKeyService from src.services.user.apikey import ApiKeyService
@@ -1203,6 +1211,14 @@ class ListAvailableModelsAdapter(AuthenticatedApiAdapter):
.all() .all()
) )
# 查询当前用户的每模型调用次数
user_usage_rows = (
db.query(UserModelUsageCount.model, UserModelUsageCount.usage_count)
.filter(UserModelUsageCount.user_id == user.id)
.all()
)
user_usage_map: dict[str, int] = {row.model: row.usage_count for row in user_usage_rows}
# 转换为响应格式(复用 PublicGlobalModelResponse schema # 转换为响应格式(复用 PublicGlobalModelResponse schema
model_responses = [ model_responses = [
PublicGlobalModelResponse( PublicGlobalModelResponse(
@@ -1214,6 +1230,7 @@ class ListAvailableModelsAdapter(AuthenticatedApiAdapter):
default_tiered_pricing=gm.default_tiered_pricing, default_tiered_pricing=gm.default_tiered_pricing,
supported_capabilities=gm.supported_capabilities, supported_capabilities=gm.supported_capabilities,
config=gm.config, config=gm.config,
usage_count=user_usage_map.get(gm.name, 0),
) )
for gm in models for gm in models
] ]

View File

@@ -100,17 +100,15 @@ def check_capability_match(
匹配逻辑: 匹配逻辑:
1. EXCLUSIVE互斥能力 1. EXCLUSIVE互斥能力
- 请求需要且 Key 有 通过 - 请求需要且 Key 有 -> 通过
- 请求需要但 Key 没有 拒绝 - 请求需要但 Key 没有 -> 拒绝
- 请求不需要但 Key 有 拒绝(避免浪费高价资源) - 请求不需要但 Key 有 -> 拒绝(避免浪费高价资源)
- 请求不需要且 Key 没有 通过 - 请求不需要且 Key 没有 -> 通过
- 请求未声明但 Key 有 拒绝(关键:未声明等同于不需要) - 请求未声明但 Key 有 -> 拒绝(关键:未声明等同于不需要)
2. COMPATIBLE兼容能力 2. COMPATIBLE兼容能力
- 请求需要且 Key 有 → 通过 - 不做硬过滤,交由排序阶段通过 compute_capability_score() 处理
- 请求需要但 Key 没有 → 拒绝 - 有能力的 Key 优先排序,没有的也不被排除
- 请求不需要/未声明且 Key 有 → 通过(无额外成本,不浪费)
- 请求不需要/未声明且 Key 没有 → 通过
Args: Args:
key_capabilities: Key 拥有的能力 {"cache_1h": True, ...} key_capabilities: Key 拥有的能力 {"cache_1h": True, ...}
@@ -136,9 +134,7 @@ def check_capability_match(
if not is_required and key_has_cap: if not is_required and key_has_cap:
return False, f"不需要{cap_def.display_name}(避免浪费高价资源)" return False, f"不需要{cap_def.display_name}(避免浪费高价资源)"
elif cap_def.match_mode == CapabilityMatchMode.COMPATIBLE: # COMPATIBLE: 不做硬过滤
if is_required and not key_has_cap:
return False, f"需要{cap_def.display_name}但 Key 不支持"
# 第二步:检查 Key 拥有的 EXCLUSIVE 能力是否被请求需要 # 第二步:检查 Key 拥有的 EXCLUSIVE 能力是否被请求需要
# 如果 Key 有某个 EXCLUSIVE 能力,但请求没有声明需要,应该跳过这个 Key # 如果 Key 有某个 EXCLUSIVE 能力,但请求没有声明需要,应该跳过这个 Key
@@ -158,6 +154,40 @@ def check_capability_match(
return True, None return True, None
def compute_capability_score(
key_capabilities: dict[str, bool] | None,
requirements: dict[str, bool] | None,
) -> int:
"""
计算 COMPATIBLE 能力不匹配数量
返回 0 表示完全匹配(或无 COMPATIBLE 需求),正数表示有 N 个 COMPATIBLE 能力不满足。
用于候选排序:得分越低越优先。
Args:
key_capabilities: Key 拥有的能力
requirements: 请求需要的能力
Returns:
不满足的 COMPATIBLE 能力数量
"""
key_caps = key_capabilities or {}
reqs = requirements or {}
miss_count = 0
for cap_name, is_required in reqs.items():
if not is_required:
continue
cap_def = _capabilities.get(cap_name)
if not cap_def:
continue
if cap_def.match_mode == CapabilityMatchMode.COMPATIBLE:
if not key_caps.get(cap_name, False):
miss_count += 1
return miss_count
def _match_error_patterns(error_msg: str, patterns: list[str]) -> bool: def _match_error_patterns(error_msg: str, patterns: list[str]) -> bool:
"""检查错误信息是否匹配模式(所有关键词都要出现)""" """检查错误信息是否匹配模式(所有关键词都要出现)"""
if not patterns: if not patterns:
@@ -220,20 +250,14 @@ class _CapabilityDefinitionsProxy:
CAPABILITY_DEFINITIONS = _CapabilityDefinitionsProxy() CAPABILITY_DEFINITIONS = _CapabilityDefinitionsProxy()
# ============ 兼容旧的插件基类(逐步废弃) ============
CapabilityPlugin = CapabilityDefinition # 类型别名,兼容旧代码
# ============ 注册内置能力 ============ # ============ 注册内置能力 ============
register_capability( register_capability(
name="cache_1h", name="cache_1h",
display_name="1 小时缓存", display_name="1 小时缓存",
description="使用 1 小时缓存 TTL价格更高适合长对话", description="使用 1 小时缓存 TTL价格更高适合长对话",
match_mode=CapabilityMatchMode.EXCLUSIVE, match_mode=CapabilityMatchMode.COMPATIBLE,
config_mode=CapabilityConfigMode.USER_CONFIGURABLE, config_mode=CapabilityConfigMode.REQUEST_PARAM,
short_name="1h缓存", short_name="1h缓存",
) )
@@ -251,7 +275,7 @@ register_capability(
name="gemini_files", name="gemini_files",
display_name="Gemini 文件 API", display_name="Gemini 文件 API",
description="支持 Gemini Files API文件上传/管理),仅 Google 官方 API 支持", description="支持 Gemini Files API文件上传/管理),仅 Google 官方 API 支持",
match_mode=CapabilityMatchMode.COMPATIBLE, match_mode=CapabilityMatchMode.EXCLUSIVE,
config_mode=CapabilityConfigMode.USER_CONFIGURABLE, config_mode=CapabilityConfigMode.REQUEST_PARAM,
short_name="文件API", short_name="文件API",
) )

View File

@@ -577,7 +577,7 @@ class ModelResponse(BaseModel):
id: str id: str
provider_id: str provider_id: str
global_model_id: str | None global_model_id: str
provider_model_name: str provider_model_name: str
provider_model_mappings: list[dict] | None = None provider_model_mappings: list[dict] | None = None
@@ -612,7 +612,7 @@ class ModelResponse(BaseModel):
created_at: datetime created_at: datetime
updated_at: datetime updated_at: datetime
# 关联的 GlobalModel 信息(如果有) # 关联的 GlobalModel 信息
global_model_name: str | None = None global_model_name: str | None = None
global_model_display_name: str | None = None global_model_display_name: str | None = None
@@ -744,6 +744,8 @@ class PublicGlobalModelResponse(BaseModel):
supported_capabilities: list[str] | None = None supported_capabilities: list[str] | None = None
# 模型配置JSON # 模型配置JSON
config: dict | None = None config: dict | None = None
# 调用次数
usage_count: int = 0
class PublicGlobalModelListResponse(BaseModel): class PublicGlobalModelListResponse(BaseModel):

View File

@@ -1001,8 +1001,8 @@ class GlobalModel(ExportMixin, Base):
# "cache_creation_price_per_1m": 3.75, # 可选 # "cache_creation_price_per_1m": 3.75, # 可选
# "cache_read_price_per_1m": 0.30, # 可选 # "cache_read_price_per_1m": 0.30, # 可选
# "cache_ttl_pricing": [ # 可选:按缓存时长分价格 # "cache_ttl_pricing": [ # 可选:按缓存时长分价格
# {"ttl_minutes": 5, "cache_read_price_per_1m": 0.30}, # {"ttl_minutes": 5, "cache_creation_price_per_1m": 3.75, "cache_read_price_per_1m": 0.30},
# {"ttl_minutes": 60, "cache_read_price_per_1m": 0.50} # {"ttl_minutes": 60, "cache_creation_price_per_1m": 6.00, "cache_read_price_per_1m": 0.50}
# ] # ]
# }, # },
# {"up_to": null, "input_price_per_1m": 1.25, ...} # {"up_to": null, "input_price_per_1m": 1.25, ...}
@@ -2700,5 +2700,36 @@ class GeminiFileMapping(Base):
) )
class UserModelUsageCount(Base):
"""用户-模型维度调用次数计数器
每个用户对每个模型维护一个原子递增的计数器,
避免从 Usage 表聚合查询,查询性能 O(N) 其中 N 是用户使用过的模型数。
"""
__tablename__ = "user_model_usage_counts"
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
user_id = Column(String(36), ForeignKey("users.id", ondelete="CASCADE"), nullable=False)
model = Column(String(100), nullable=False)
usage_count = Column(Integer, default=0, nullable=False)
created_at = Column(
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
)
updated_at = Column(
DateTime(timezone=True),
default=lambda: datetime.now(timezone.utc),
onupdate=lambda: datetime.now(timezone.utc),
nullable=False,
)
__table_args__ = (
UniqueConstraint("user_id", "model", name="uq_user_model_usage_count"),
Index("idx_user_model_usage_user", "user_id"),
Index("idx_user_model_usage_model", "model"),
)
# 导入扩展的数据库模型 # 导入扩展的数据库模型
from .database_extensions import ApiKeyProviderMapping, ProviderUsageTracking from .database_extensions import ApiKeyProviderMapping, ProviderUsageTracking

View File

@@ -69,8 +69,8 @@ class GlobalModel(ExportMixin, Base):
# "cache_creation_price_per_1m": 3.75, # 可选 # "cache_creation_price_per_1m": 3.75, # 可选
# "cache_read_price_per_1m": 0.30, # 可选 # "cache_read_price_per_1m": 0.30, # 可选
# "cache_ttl_pricing": [ # 可选:按缓存时长分价格 # "cache_ttl_pricing": [ # 可选:按缓存时长分价格
# {"ttl_minutes": 5, "cache_read_price_per_1m": 0.30}, # {"ttl_minutes": 5, "cache_creation_price_per_1m": 3.75, "cache_read_price_per_1m": 0.30},
# {"ttl_minutes": 60, "cache_read_price_per_1m": 0.50} # {"ttl_minutes": 60, "cache_creation_price_per_1m": 6.00, "cache_read_price_per_1m": 0.50}
# ] # ]
# }, # },
# {"up_to": null, "input_price_per_1m": 1.25, ...} # {"up_to": null, "input_price_per_1m": 1.25, ...}
@@ -132,9 +132,7 @@ class Model(ExportMixin, Base):
设计原则: 设计原则:
- Model 表示 Provider 对某个模型的具体实现 - Model 表示 Provider 对某个模型的具体实现
- global_model_id 可为空: - global_model_id 必填,必须关联到一个 GlobalModel
- 为空时:模型尚未关联到 GlobalModel不参与路由
- 不为空时:模型已关联 GlobalModel参与路由
- provider_model_name 是 Provider 侧的实际模型名称 (可能与 GlobalModel.name 不同) - provider_model_name 是 Provider 侧的实际模型名称 (可能与 GlobalModel.name 不同)
- 价格和能力配置可为空,为空时使用 GlobalModel 的默认值 - 价格和能力配置可为空,为空时使用 GlobalModel 的默认值
""" """
@@ -154,8 +152,8 @@ class Model(ExportMixin, Base):
id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True) id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()), index=True)
provider_id = Column(String(36), ForeignKey("providers.id"), nullable=False) provider_id = Column(String(36), ForeignKey("providers.id"), nullable=False)
# 可为空NULL 表示未关联,不参与路由;非 NULL 表示已关联,参与路由 # 必须关联一个 GlobalModel
global_model_id = Column(String(36), ForeignKey("global_models.id"), nullable=True, index=True) global_model_id = Column(String(36), ForeignKey("global_models.id"), nullable=False, index=True)
# Provider 映射配置 # Provider 映射配置
provider_model_name = Column(String(200), nullable=False) # Provider 侧的主模型名称 provider_model_name = Column(String(200), nullable=False) # Provider 侧的主模型名称

View File

@@ -19,6 +19,9 @@ class CacheTTLPricing(BaseModel):
cache_creation_price_per_1m: float = Field( cache_creation_price_per_1m: float = Field(
..., ge=0, description="该时长的缓存创建价格/M tokens" ..., ge=0, description="该时长的缓存创建价格/M tokens"
) )
cache_read_price_per_1m: float | None = Field(
None, ge=0, description="该时长的缓存读取价格/M tokens"
)
class PricingTier(BaseModel): class PricingTier(BaseModel):
@@ -313,8 +316,8 @@ class ImportFromUpstreamSuccessItem(BaseModel):
model_id: str = Field(..., description="上游模型 ID") model_id: str = Field(..., description="上游模型 ID")
provider_model_id: str = Field(..., description="Provider Model ID") provider_model_id: str = Field(..., description="Provider Model ID")
global_model_id: str | None = Field("", description="GlobalModel ID(如果已关联)") global_model_id: str = Field(..., description="GlobalModel ID")
global_model_name: str | None = Field("", description="GlobalModel 名称(如果已关联)") global_model_name: str = Field(..., description="GlobalModel 名称")
created_global_model: bool = Field( created_global_model: bool = Field(
False, description="是否新创建了 GlobalModel始终为 false" False, description="是否新创建了 GlobalModel始终为 false"
) )

View File

@@ -71,6 +71,8 @@ class Usage(Base):
# 缓存相关 tokens (for Claude models) # 缓存相关 tokens (for Claude models)
cache_creation_input_tokens = Column(Integer, default=0) cache_creation_input_tokens = Column(Integer, default=0)
cache_read_input_tokens = Column(Integer, default=0) cache_read_input_tokens = Column(Integer, default=0)
cache_creation_input_tokens_5m = Column(Integer, default=0) # 5min TTL 缓存创建
cache_creation_input_tokens_1h = Column(Integer, default=0) # 1h TTL 缓存创建
# 成本计算 # 成本计算
input_cost_usd = Column(Float, default=0.0) input_cost_usd = Column(Float, default=0.0)

View File

@@ -231,7 +231,14 @@ class DefaultBillingRuleGenerator:
"source": "tiered", "source": "tiered",
"tier_key": tier_key, "tier_key": tier_key,
"allow_zero": True, "allow_zero": True,
"tiers": _tiers_for("cache_creation_price_per_1m", default_multiplier=1.25), # TTL override supported when dims include cache_ttl_minutes
"ttl_key": "cache_ttl_minutes",
"ttl_value_key": "cache_creation_price_per_1m",
"tiers": _tiers_for(
"cache_creation_price_per_1m",
default_multiplier=1.25,
include_cache_ttl_pricing=True,
),
"default": base_cache_creation_price, "default": base_cache_creation_price,
} }
dimension_mappings["cache_read_price_per_1m"] = { dimension_mappings["cache_read_price_per_1m"] = {

View File

@@ -545,8 +545,8 @@ class GlobalModelService:
models_to_delete: list[Model] = [] models_to_delete: list[Model] = []
for model in models: for model in models:
# 跳过没有关联 GlobalModel 的 # 跳过 global_model 关系未加载
if not model.global_model_id or not model.global_model: if not model.global_model:
continue continue
global_model = cast(GlobalModel, model.global_model) global_model = cast(GlobalModel, model.global_model)

View File

@@ -65,15 +65,14 @@ class ModelService:
db.commit() db.commit()
db.refresh(model) db.refresh(model)
# 显式加载 global_model 关系 # 显式加载 global_model 关系
if model.global_model_id: from sqlalchemy.orm import joinedload
from sqlalchemy.orm import joinedload
model = ( model = (
db.query(Model) db.query(Model)
.options(joinedload(Model.global_model)) .options(joinedload(Model.global_model))
.filter(Model.id == model.id) .filter(Model.id == model.id)
.first() .first()
) )
logger.info( logger.info(
f"创建模型成功: provider={provider.name}, model={model.provider_model_name}, global_model_id={model.global_model_id}" f"创建模型成功: provider={provider.name}, model={model.provider_model_name}, global_model_id={model.global_model_id}"
@@ -226,7 +225,7 @@ class ModelService:
) )
# 清除内存缓存ModelMapperMiddleware 实例) # 清除内存缓存ModelMapperMiddleware 实例)
if model.provider_id and model.global_model_id: if model.provider_id:
cache_service = get_cache_invalidation_service() cache_service = get_cache_invalidation_service()
cache_service.on_model_changed(model.provider_id, model.global_model_id) cache_service.on_model_changed(model.provider_id, model.global_model_id)
@@ -297,7 +296,7 @@ class ModelService:
) )
# 清除内存缓存 # 清除内存缓存
if cache_info["provider_id"] and cache_info["global_model_id"]: if cache_info["provider_id"]:
cache_service = get_cache_invalidation_service() cache_service = get_cache_invalidation_service()
cache_service.on_model_changed( cache_service.on_model_changed(
cache_info["provider_id"], cache_info["global_model_id"] cache_info["provider_id"], cache_info["global_model_id"]
@@ -338,7 +337,7 @@ class ModelService:
) )
# 清除内存缓存ModelMapperMiddleware 实例) # 清除内存缓存ModelMapperMiddleware 实例)
if model.provider_id and model.global_model_id: if model.provider_id:
cache_service = get_cache_invalidation_service() cache_service = get_cache_invalidation_service()
cache_service.on_model_changed(model.provider_id, model.global_model_id) cache_service.on_model_changed(model.provider_id, model.global_model_id)

View File

@@ -19,7 +19,12 @@ from sqlalchemy.orm import Session, selectinload
from src.core.api_format.conversion.compatibility import is_format_compatible from src.core.api_format.conversion.compatibility import is_format_compatible
from src.core.api_format.enums import EndpointKind from src.core.api_format.enums import EndpointKind
from src.core.api_format.signature import make_signature_key, parse_signature_key from src.core.api_format.signature import make_signature_key, parse_signature_key
from src.core.key_capabilities import check_capability_match from src.core.key_capabilities import (
CapabilityMatchMode,
check_capability_match,
compute_capability_score,
get_capability,
)
from src.core.logger import logger from src.core.logger import logger
from src.core.model_permissions import check_model_allowed_with_mappings from src.core.model_permissions import check_model_allowed_with_mappings
from src.models.database import ( from src.models.database import (
@@ -219,9 +224,13 @@ class CandidateBuilder:
# 检查模型是否支持所需的能力(在 Provider 级别检查,而不是 Key 级别) # 检查模型是否支持所需的能力(在 Provider 级别检查,而不是 Key 级别)
# 只有当 model_supported_capabilities 非空时才进行检查 # 只有当 model_supported_capabilities 非空时才进行检查
# 空列表意味着模型没有配置能力限制,默认支持所有能力 # 空列表意味着模型没有配置能力限制,默认支持所有能力
# COMPATIBLE 能力跳过模型级硬过滤(交由排序阶段处理)
if capability_requirements and model_supported_capabilities: if capability_requirements and model_supported_capabilities:
for cap_name, is_required in capability_requirements.items(): for cap_name, is_required in capability_requirements.items():
if is_required and cap_name not in model_supported_capabilities: if is_required and cap_name not in model_supported_capabilities:
cap_def = get_capability(cap_name)
if cap_def and cap_def.match_mode == CapabilityMatchMode.COMPATIBLE:
continue
return ( return (
False, False,
f"模型 {model_name} 不支持能力: {cap_name}", f"模型 {model_name} 不支持能力: {cap_name}",
@@ -617,6 +626,15 @@ class CandidateBuilder:
needs_conversion=needs_conversion, needs_conversion=needs_conversion,
provider_api_format=str(endpoint_format_str or ""), provider_api_format=str(endpoint_format_str or ""),
output_limit=output_limit, output_limit=output_limit,
# is_skipped 候选不参与排序miss_count 无意义,置 0 避免干扰
capability_miss_count=(
compute_capability_score(
key.capabilities or {},
capability_requirements,
)
if is_available
else 0
),
) )
if needs_conversion: if needs_conversion:

View File

@@ -18,6 +18,8 @@ from src.services.scheduling.utils import affinity_hash
from src.services.system.config import SystemConfigService from src.services.system.config import SystemConfigService
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import Callable
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from src.models.database import ProviderAPIKey from src.models.database import ProviderAPIKey
@@ -30,6 +32,28 @@ class CandidateSorter:
def __init__(self, config: SchedulingConfig) -> None: def __init__(self, config: SchedulingConfig) -> None:
self._config = config self._config = config
@staticmethod
def _split_by_capability_match(
candidates: list[ProviderCandidate],
) -> tuple[list[ProviderCandidate], list[ProviderCandidate]]:
"""按 capability_miss_count 分组:完全匹配(0)在前,部分匹配(>0)在后"""
full_match = [c for c in candidates if c.capability_miss_count == 0]
partial_match = [c for c in candidates if c.capability_miss_count > 0]
return full_match, partial_match
def _with_capability_split(
self,
candidates: list[ProviderCandidate],
sort_fn: Callable[..., list[ProviderCandidate]],
*args: object,
**kwargs: object,
) -> list[ProviderCandidate]:
"""通用包装:先按 capability_miss_count 分组,再分别排序后合并"""
if not candidates:
return candidates
full_match, partial_match = self._split_by_capability_match(candidates)
return sort_fn(full_match, *args, **kwargs) + sort_fn(partial_match, *args, **kwargs)
def _apply_priority_mode_sort( def _apply_priority_mode_sort(
self, self,
candidates: list[ProviderCandidate], candidates: list[ProviderCandidate],
@@ -40,7 +64,8 @@ class CandidateSorter:
""" """
根据优先级模式对候选列表排序(数字越小越优先) 根据优先级模式对候选列表排序(数字越小越优先)
排序规则(受 keep_priority_on_conversion 配置影响) 排序规则:
0. 按 capability_miss_count 分组:完全匹配(0)在前,部分匹配(>0)在后
1. 如果全局配置 keep_priority_on_conversion=True所有候选保持原优先级 1. 如果全局配置 keep_priority_on_conversion=True所有候选保持原优先级
2. 否则,按 needs_conversion 和 provider.keep_priority_on_conversion 分组: 2. 否则,按 needs_conversion 和 provider.keep_priority_on_conversion 分组:
- 保持优先级的候选exact 或 provider.keep_priority_on_conversion=True按原优先级排序 - 保持优先级的候选exact 或 provider.keep_priority_on_conversion=True按原优先级排序
@@ -52,6 +77,21 @@ class CandidateSorter:
if not candidates: if not candidates:
return candidates return candidates
return self._with_capability_split(
candidates, self._apply_priority_mode_sort_inner, db, affinity_key, api_format
)
def _apply_priority_mode_sort_inner(
self,
candidates: list[ProviderCandidate],
db: Session,
affinity_key: str | None = None,
api_format: str | None = None,
) -> list[ProviderCandidate]:
"""优先级模式排序的内部实现(不含 capability_miss_count 分组)"""
if not candidates:
return candidates
# 全局配置:如果开启,所有候选保持原优先级 # 全局配置:如果开启,所有候选保持原优先级
global_keep_priority = SystemConfigService.is_keep_priority_on_conversion(db) global_keep_priority = SystemConfigService.is_keep_priority_on_conversion(db)
@@ -159,6 +199,7 @@ class CandidateSorter:
负载均衡模式:同优先级内随机轮换 负载均衡模式:同优先级内随机轮换
排序逻辑: 排序逻辑:
0. 按 capability_miss_count 分组:完全匹配(0)在前,部分匹配(>0)在后
1. 按优先级分组provider_priority, internal_priority 或 global_priority_by_format 1. 按优先级分组provider_priority, internal_priority 或 global_priority_by_format
2. 同优先级组内随机打乱 2. 同优先级组内随机打乱
3. 不考虑缓存亲和性 3. 不考虑缓存亲和性
@@ -166,6 +207,15 @@ class CandidateSorter:
if not candidates: if not candidates:
return candidates return candidates
return self._with_capability_split(candidates, self._apply_load_balance_inner, api_format)
def _apply_load_balance_inner(
self, candidates: list[ProviderCandidate], api_format: str | None = None
) -> list[ProviderCandidate]:
"""负载均衡排序的内部实现(不含 capability_miss_count 分组)"""
if not candidates:
return candidates
priority_groups: dict[tuple, list[ProviderCandidate]] = defaultdict(list) priority_groups: dict[tuple, list[ProviderCandidate]] = defaultdict(list)
# 根据优先级模式选择分组方式 # 根据优先级模式选择分组方式

View File

@@ -29,6 +29,7 @@ class ProviderCandidate:
needs_conversion: bool = False # 是否需要格式转换 needs_conversion: bool = False # 是否需要格式转换
provider_api_format: str = "" # Provider 端点实际格式(用于健康度/熔断 bucket provider_api_format: str = "" # Provider 端点实际格式(用于健康度/熔断 bucket
output_limit: int | None = None # GlobalModel 配置的模型输出上限 output_limit: int | None = None # GlobalModel 配置的模型输出上限
capability_miss_count: int = 0 # COMPATIBLE 能力不匹配数0=完全匹配,用于排序)
def _stable_order_key(self) -> tuple[int, int, str, str, str]: def _stable_order_key(self) -> tuple[int, int, str, str, str]:
""" """

View File

@@ -133,6 +133,8 @@ class UsageBillingIntegrationMixin:
output_tokens=params.output_tokens, output_tokens=params.output_tokens,
cache_creation_input_tokens=params.cache_creation_input_tokens, cache_creation_input_tokens=params.cache_creation_input_tokens,
cache_read_input_tokens=params.cache_read_input_tokens, cache_read_input_tokens=params.cache_read_input_tokens,
cache_creation_input_tokens_5m=params.cache_creation_input_tokens_5m,
cache_creation_input_tokens_1h=params.cache_creation_input_tokens_1h,
request_type=params.request_type, request_type=params.request_type,
api_format=params.api_format, api_format=params.api_format,
api_family=params.api_family, api_family=params.api_family,

View File

@@ -55,6 +55,8 @@ def build_usage_params(
output_tokens: int, output_tokens: int,
cache_creation_input_tokens: int, cache_creation_input_tokens: int,
cache_read_input_tokens: int, cache_read_input_tokens: int,
cache_creation_input_tokens_5m: int = 0,
cache_creation_input_tokens_1h: int = 0,
request_type: str, request_type: str,
api_format: str | None, api_format: str | None,
api_family: str | None = None, api_family: str | None = None,
@@ -201,6 +203,8 @@ def build_usage_params(
"total_tokens": input_tokens + output_tokens, "total_tokens": input_tokens + output_tokens,
"cache_creation_input_tokens": cache_creation_input_tokens, "cache_creation_input_tokens": cache_creation_input_tokens,
"cache_read_input_tokens": cache_read_input_tokens, "cache_read_input_tokens": cache_read_input_tokens,
"cache_creation_input_tokens_5m": cache_creation_input_tokens_5m,
"cache_creation_input_tokens_1h": cache_creation_input_tokens_1h,
"input_cost_usd": input_cost, "input_cost_usd": input_cost,
"output_cost_usd": output_cost, "output_cost_usd": output_cost,
"cache_cost_usd": cache_cost, "cache_cost_usd": cache_cost,
@@ -292,6 +296,12 @@ def update_existing_usage(
existing_usage.total_tokens = usage_params["total_tokens"] existing_usage.total_tokens = usage_params["total_tokens"]
existing_usage.cache_creation_input_tokens = usage_params["cache_creation_input_tokens"] existing_usage.cache_creation_input_tokens = usage_params["cache_creation_input_tokens"]
existing_usage.cache_read_input_tokens = usage_params["cache_read_input_tokens"] existing_usage.cache_read_input_tokens = usage_params["cache_read_input_tokens"]
existing_usage.cache_creation_input_tokens_5m = usage_params.get(
"cache_creation_input_tokens_5m", 0
)
existing_usage.cache_creation_input_tokens_1h = usage_params.get(
"cache_creation_input_tokens_1h", 0
)
existing_usage.input_cost_usd = usage_params["input_cost_usd"] existing_usage.input_cost_usd = usage_params["input_cost_usd"]
existing_usage.output_cost_usd = usage_params["output_cost_usd"] existing_usage.output_cost_usd = usage_params["output_cost_usd"]
existing_usage.cache_cost_usd = usage_params["cache_cost_usd"] existing_usage.cache_cost_usd = usage_params["cache_cost_usd"]

View File

@@ -49,6 +49,8 @@ class UsageRecordParams:
cache_ttl_minutes: int | None cache_ttl_minutes: int | None
use_tiered_pricing: bool use_tiered_pricing: bool
target_model: str | None target_model: str | None
cache_creation_input_tokens_5m: int = 0
cache_creation_input_tokens_1h: int = 0
def __post_init__(self) -> None: def __post_init__(self) -> None:
"""验证关键字段,确保数据完整性""" """验证关键字段,确保数据完整性"""

View File

@@ -68,6 +68,8 @@ def _event_to_record(event: UsageEvent) -> dict[str, Any]:
"output_tokens": data.get("output_tokens") or 0, "output_tokens": data.get("output_tokens") or 0,
"cache_creation_input_tokens": data.get("cache_creation_input_tokens") or 0, "cache_creation_input_tokens": data.get("cache_creation_input_tokens") or 0,
"cache_read_input_tokens": data.get("cache_read_input_tokens") or 0, "cache_read_input_tokens": data.get("cache_read_input_tokens") or 0,
"cache_creation_input_tokens_5m": data.get("cache_creation_input_tokens_5m") or 0,
"cache_creation_input_tokens_1h": data.get("cache_creation_input_tokens_1h") or 0,
"request_type": data.get("request_type") or "chat", "request_type": data.get("request_type") or "chat",
"api_format": data.get("api_format"), "api_format": data.get("api_format"),
"api_family": data.get("api_family"), "api_family": data.get("api_family"),

View File

@@ -7,7 +7,7 @@ from typing import Any
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from src.core.logger import logger from src.core.logger import logger
from src.models.database import ApiKey, Provider, Usage, User from src.models.database import ApiKey, Provider, Usage, User, UserModelUsageCount
from src.services.usage._billing_integration import UsageBillingIntegrationMixin from src.services.usage._billing_integration import UsageBillingIntegrationMixin
from src.services.usage._recording_helpers import ( from src.services.usage._recording_helpers import (
METADATA_KEEP_KEYS, METADATA_KEEP_KEYS,
@@ -44,6 +44,31 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
"""更新已存在的 Usage 记录(委托到模块级函数)""" """更新已存在的 Usage 记录(委托到模块级函数)"""
update_existing_usage(existing_usage, usage_params, target_model) update_existing_usage(existing_usage, usage_params, target_model)
@staticmethod
def _increment_user_model_usage(
db: Session, user: User | None, model: str, count: int = 1
) -> None:
"""原子递增用户-模型调用次数计数器"""
if user is None:
return
from sqlalchemy import func as sa_func
from sqlalchemy.dialects.postgresql import insert as pg_insert
stmt = pg_insert(UserModelUsageCount).values(
id=str(uuid.uuid4()),
user_id=user.id,
model=model,
usage_count=count,
)
stmt = stmt.on_conflict_do_update(
constraint="uq_user_model_usage_count",
set_={
"usage_count": UserModelUsageCount.usage_count + count,
"updated_at": sa_func.now(),
},
)
db.execute(stmt)
@classmethod @classmethod
def _sanitize_request_metadata(cls, metadata: dict[str, Any]) -> dict[str, Any]: def _sanitize_request_metadata(cls, metadata: dict[str, Any]) -> dict[str, Any]:
"""元数据清理(委托到模块级函数)""" """元数据清理(委托到模块级函数)"""
@@ -65,6 +90,8 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
output_tokens: int, output_tokens: int,
cache_creation_input_tokens: int = 0, cache_creation_input_tokens: int = 0,
cache_read_input_tokens: int = 0, cache_read_input_tokens: int = 0,
cache_creation_input_tokens_5m: int = 0,
cache_creation_input_tokens_1h: int = 0,
request_type: str = "chat", request_type: str = "chat",
api_format: str | None = None, api_format: str | None = None,
api_family: str | None = None, api_family: str | None = None,
@@ -116,6 +143,8 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
output_tokens=output_tokens, output_tokens=output_tokens,
cache_creation_input_tokens=cache_creation_input_tokens, cache_creation_input_tokens=cache_creation_input_tokens,
cache_read_input_tokens=cache_read_input_tokens, cache_read_input_tokens=cache_read_input_tokens,
cache_creation_input_tokens_5m=cache_creation_input_tokens_5m,
cache_creation_input_tokens_1h=cache_creation_input_tokens_1h,
request_type=request_type, request_type=request_type,
api_format=api_format, api_format=api_format,
api_family=api_family, api_family=api_family,
@@ -162,6 +191,9 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
.values(usage_count=GlobalModel.usage_count + 1) .values(usage_count=GlobalModel.usage_count + 1)
) )
# 更新用户-模型调用次数计数器
cls._increment_user_model_usage(db, user, model)
# 更新 Provider 月度使用量(原子操作) # 更新 Provider 月度使用量(原子操作)
if provider_id: if provider_id:
actual_total_cost = usage_params["actual_total_cost_usd"] actual_total_cost = usage_params["actual_total_cost_usd"]
@@ -191,6 +223,8 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
output_tokens: int, output_tokens: int,
cache_creation_input_tokens: int = 0, cache_creation_input_tokens: int = 0,
cache_read_input_tokens: int = 0, cache_read_input_tokens: int = 0,
cache_creation_input_tokens_5m: int = 0,
cache_creation_input_tokens_1h: int = 0,
request_type: str = "chat", request_type: str = "chat",
api_format: str | None = None, api_format: str | None = None,
api_family: str | None = None, api_family: str | None = None,
@@ -244,6 +278,8 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
output_tokens=output_tokens, output_tokens=output_tokens,
cache_creation_input_tokens=cache_creation_input_tokens, cache_creation_input_tokens=cache_creation_input_tokens,
cache_read_input_tokens=cache_read_input_tokens, cache_read_input_tokens=cache_read_input_tokens,
cache_creation_input_tokens_5m=cache_creation_input_tokens_5m,
cache_creation_input_tokens_1h=cache_creation_input_tokens_1h,
request_type=request_type, request_type=request_type,
api_format=api_format, api_format=api_format,
api_family=api_family, api_family=api_family,
@@ -347,6 +383,9 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
.values(usage_count=GlobalModel.usage_count + 1) .values(usage_count=GlobalModel.usage_count + 1)
) )
# 更新用户-模型调用次数计数器
cls._increment_user_model_usage(db, user, model)
# 更新 Provider 月度使用量 # 更新 Provider 月度使用量
if provider_id: if provider_id:
actual_total_cost = usage_params["actual_total_cost_usd"] actual_total_cost = usage_params["actual_total_cost_usd"]
@@ -387,6 +426,8 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
output_tokens: int = 0, output_tokens: int = 0,
cache_creation_input_tokens: int = 0, cache_creation_input_tokens: int = 0,
cache_read_input_tokens: int = 0, cache_read_input_tokens: int = 0,
cache_creation_input_tokens_5m: int = 0,
cache_creation_input_tokens_1h: int = 0,
api_format: str | None = None, api_format: str | None = None,
api_family: str | None = None, api_family: str | None = None,
endpoint_kind: str | None = None, endpoint_kind: str | None = None,
@@ -451,6 +492,8 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
output_tokens=output_tokens, output_tokens=output_tokens,
cache_creation_input_tokens=cache_creation_input_tokens, cache_creation_input_tokens=cache_creation_input_tokens,
cache_read_input_tokens=cache_read_input_tokens, cache_read_input_tokens=cache_read_input_tokens,
cache_creation_input_tokens_5m=cache_creation_input_tokens_5m,
cache_creation_input_tokens_1h=cache_creation_input_tokens_1h,
request_type=request_type, request_type=request_type,
api_format=api_format, api_format=api_format,
api_family=api_family, api_family=api_family,
@@ -588,6 +631,9 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
.values(usage_count=GlobalModel.usage_count + 1) .values(usage_count=GlobalModel.usage_count + 1)
) )
# 更新用户-模型调用次数计数器
cls._increment_user_model_usage(db, user, model)
# 更新 Provider 月度使用量(使用 actual_total_cost # 更新 Provider 月度使用量(使用 actual_total_cost
if provider_id: if provider_id:
actual_total_cost = usage_params["actual_total_cost_usd"] actual_total_cost = usage_params["actual_total_cost_usd"]
@@ -714,6 +760,9 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
lambda: {"requests": 0, "cost": 0.0, "is_standalone": False} lambda: {"requests": 0, "cost": 0.0, "is_standalone": False}
) )
model_counts: dict[str, int] = defaultdict(int) # model -> count model_counts: dict[str, int] = defaultdict(int) # model -> count
user_model_counts: dict[tuple[str, str], int] = defaultdict(
int
) # (user_id, model) -> count
provider_costs: dict[str, float] = defaultdict(float) # provider_id -> cost provider_costs: dict[str, float] = defaultdict(float) # provider_id -> cost
# 合并所有需要处理的记录(用于预取 user/api_key # 合并所有需要处理的记录(用于预取 user/api_key
@@ -754,6 +803,12 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
output_tokens=int(record.get("output_tokens") or 0), output_tokens=int(record.get("output_tokens") or 0),
cache_creation_input_tokens=int(record.get("cache_creation_input_tokens") or 0), cache_creation_input_tokens=int(record.get("cache_creation_input_tokens") or 0),
cache_read_input_tokens=int(record.get("cache_read_input_tokens") or 0), cache_read_input_tokens=int(record.get("cache_read_input_tokens") or 0),
cache_creation_input_tokens_5m=int(
record.get("cache_creation_input_tokens_5m") or 0
),
cache_creation_input_tokens_1h=int(
record.get("cache_creation_input_tokens_1h") or 0
),
request_type=record.get("request_type") or "chat", request_type=record.get("request_type") or "chat",
api_format=record.get("api_format"), api_format=record.get("api_format"),
api_family=record.get("api_family"), api_family=record.get("api_family"),
@@ -850,6 +905,8 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
# 聚合统计 # 聚合统计
model_name = record.get("model") or "unknown" model_name = record.get("model") or "unknown"
model_counts[model_name] += 1 model_counts[model_name] += 1
if user:
user_model_counts[(str(user.id), model_name)] += 1
provider_id = record.get("provider_id") provider_id = record.get("provider_id")
if provider_id: if provider_id:
@@ -898,6 +955,8 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
# 聚合统计 # 聚合统计
model_name = record.get("model") or "unknown" model_name = record.get("model") or "unknown"
model_counts[model_name] += 1 model_counts[model_name] += 1
if user:
user_model_counts[(str(user.id), model_name)] += 1
provider_id = record.get("provider_id") provider_id = record.get("provider_id")
if provider_id: if provider_id:
@@ -959,6 +1018,30 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
.values(usage_count=GlobalModel.usage_count + count) .values(usage_count=GlobalModel.usage_count + count)
) )
# 批量更新用户-模型调用次数计数器
from sqlalchemy import func as sql_func
from sqlalchemy.dialects.postgresql import insert as pg_insert
if user_model_counts:
rows = [
{
"id": str(uuid.uuid4()),
"user_id": uid,
"model": model_name,
"usage_count": count,
}
for (uid, model_name), count in user_model_counts.items()
]
stmt = pg_insert(UserModelUsageCount).values(rows)
stmt = stmt.on_conflict_do_update(
constraint="uq_user_model_usage_count",
set_={
"usage_count": UserModelUsageCount.usage_count + stmt.excluded.usage_count,
"updated_at": sql_func.now(),
},
)
db.execute(stmt)
# 批量更新 Provider 月度使用量 # 批量更新 Provider 月度使用量
for provider_id, cost in provider_costs.items(): for provider_id, cost in provider_costs.items():
if cost > 0: if cost > 0:
@@ -969,8 +1052,6 @@ class UsageRecordingMixin(UsageBillingIntegrationMixin):
) )
# 批量更新用户使用量 # 批量更新用户使用量
from sqlalchemy import func as sql_func
for user_id, cost in user_costs.items(): for user_id, cost in user_costs.items():
if cost > 0: if cost > 0:
db.execute( db.execute(

View File

@@ -100,6 +100,8 @@ class StreamUsageTracker:
self.output_tokens = 0 self.output_tokens = 0
self.cache_creation_input_tokens = 0 self.cache_creation_input_tokens = 0
self.cache_read_input_tokens = 0 self.cache_read_input_tokens = 0
self.cache_creation_input_tokens_5m = 0
self.cache_creation_input_tokens_1h = 0
self.accumulated_content = "" self.accumulated_content = ""
# 完整响应跟踪(仅用于内部统计,不记录到数据库) # 完整响应跟踪(仅用于内部统计,不记录到数据库)
@@ -477,6 +479,8 @@ class StreamUsageTracker:
""" """
import time import time
from src.api.handlers.base.utils import extract_cache_creation_tokens_detail
self.start_time = time.time() self.start_time = time.time()
self.request_data = request_data # 保存请求数据 self.request_data = request_data # 保存请求数据
@@ -545,21 +549,17 @@ class StreamUsageTracker:
# 如果响应中包含准确的usage信息使用它 # 如果响应中包含准确的usage信息使用它
self.input_tokens = usage.get("input_tokens", self.input_tokens) self.input_tokens = usage.get("input_tokens", self.input_tokens)
self.output_tokens = usage.get("output_tokens", self.output_tokens) self.output_tokens = usage.get("output_tokens", self.output_tokens)
self.cache_creation_input_tokens = usage.get(
"cache_creation_input_tokens", self.cache_creation_input_tokens
)
self.cache_read_input_tokens = usage.get( self.cache_read_input_tokens = usage.get(
"cache_read_input_tokens", self.cache_read_input_tokens "cache_read_input_tokens", self.cache_read_input_tokens
) )
# 处理新的cache_creation格式 # 统一提取 cache_creation tokens新格式优先于旧格式
if "cache_creation" in usage: total, t5m, t1h = extract_cache_creation_tokens_detail(usage)
cache_creation_data = usage.get("cache_creation", {}) if total:
# 如果没有cache_creation_input_tokens尝试从cache_creation中获取 self.cache_creation_input_tokens = total
if not self.cache_creation_input_tokens: if t5m or t1h:
self.cache_creation_input_tokens = cache_creation_data.get( self.cache_creation_input_tokens_5m = t5m
"ephemeral_5m_input_tokens", 0 self.cache_creation_input_tokens_1h = t1h
) + cache_creation_data.get("ephemeral_1h_input_tokens", 0)
finally: finally:
# 流结束后记录使用量 # 流结束后记录使用量
@@ -768,6 +768,8 @@ class StreamUsageTracker:
output_tokens=self.output_tokens, output_tokens=self.output_tokens,
cache_creation_input_tokens=self.cache_creation_input_tokens, cache_creation_input_tokens=self.cache_creation_input_tokens,
cache_read_input_tokens=self.cache_read_input_tokens, cache_read_input_tokens=self.cache_read_input_tokens,
cache_creation_input_tokens_5m=self.cache_creation_input_tokens_5m,
cache_creation_input_tokens_1h=self.cache_creation_input_tokens_1h,
request_type="chat", request_type="chat",
api_format=self.api_format, api_format=self.api_format,
api_family=self.api_family, api_family=self.api_family,

View File

@@ -70,6 +70,8 @@ class MessageTelemetry:
client_response_headers: dict[str, Any] | None = None, client_response_headers: dict[str, Any] | None = None,
cache_creation_tokens: int = 0, cache_creation_tokens: int = 0,
cache_read_tokens: int = 0, cache_read_tokens: int = 0,
cache_creation_tokens_5m: int = 0,
cache_creation_tokens_1h: int = 0,
is_stream: bool = False, is_stream: bool = False,
provider_request_headers: dict[str, Any] | None = None, provider_request_headers: dict[str, Any] | None = None,
provider_request_body: Any | None = None, provider_request_body: Any | None = None,
@@ -111,6 +113,8 @@ class MessageTelemetry:
output_tokens=output_tokens, output_tokens=output_tokens,
cache_creation_input_tokens=cache_creation_tokens, cache_creation_input_tokens=cache_creation_tokens,
cache_read_input_tokens=cache_read_tokens, cache_read_input_tokens=cache_read_tokens,
cache_creation_input_tokens_5m=cache_creation_tokens_5m,
cache_creation_input_tokens_1h=cache_creation_tokens_1h,
request_type="chat", request_type="chat",
api_format=api_format, api_format=api_format,
api_family=api_family, api_family=api_family,
@@ -181,6 +185,8 @@ class MessageTelemetry:
output_tokens: int = 0, output_tokens: int = 0,
cache_creation_tokens: int = 0, cache_creation_tokens: int = 0,
cache_read_tokens: int = 0, cache_read_tokens: int = 0,
cache_creation_tokens_5m: int = 0,
cache_creation_tokens_1h: int = 0,
response_body: dict[str, Any] | None = None, response_body: dict[str, Any] | None = None,
response_headers: dict[str, Any] | None = None, response_headers: dict[str, Any] | None = None,
client_response_headers: dict[str, Any] | None = None, client_response_headers: dict[str, Any] | None = None,
@@ -227,6 +233,8 @@ class MessageTelemetry:
output_tokens=output_tokens, output_tokens=output_tokens,
cache_creation_input_tokens=cache_creation_tokens, cache_creation_input_tokens=cache_creation_tokens,
cache_read_input_tokens=cache_read_tokens, cache_read_input_tokens=cache_read_tokens,
cache_creation_input_tokens_5m=cache_creation_tokens_5m,
cache_creation_input_tokens_1h=cache_creation_tokens_1h,
request_type="chat", request_type="chat",
api_format=api_format, api_format=api_format,
api_family=api_family, api_family=api_family,
@@ -276,6 +284,8 @@ class MessageTelemetry:
output_tokens: int = 0, output_tokens: int = 0,
cache_creation_tokens: int = 0, cache_creation_tokens: int = 0,
cache_read_tokens: int = 0, cache_read_tokens: int = 0,
cache_creation_tokens_5m: int = 0,
cache_creation_tokens_1h: int = 0,
response_body: dict[str, Any] | None = None, response_body: dict[str, Any] | None = None,
response_headers: dict[str, Any] | None = None, response_headers: dict[str, Any] | None = None,
client_response_headers: dict[str, Any] | None = None, client_response_headers: dict[str, Any] | None = None,
@@ -308,6 +318,8 @@ class MessageTelemetry:
output_tokens=output_tokens, output_tokens=output_tokens,
cache_creation_input_tokens=cache_creation_tokens, cache_creation_input_tokens=cache_creation_tokens,
cache_read_input_tokens=cache_read_tokens, cache_read_input_tokens=cache_read_tokens,
cache_creation_input_tokens_5m=cache_creation_tokens_5m,
cache_creation_input_tokens_1h=cache_creation_tokens_1h,
request_type="chat", request_type="chat",
api_format=api_format, api_format=api_format,
api_family=api_family, api_family=api_family,

View File

@@ -206,6 +206,14 @@ class QueueTelemetryWriter(TelemetryWriter):
if cache_read: if cache_read:
data["cache_read_input_tokens"] = cache_read data["cache_read_input_tokens"] = cache_read
# 缓存 5m/1h 细分
cache_creation_5m = kwargs.get("cache_creation_tokens_5m", 0)
cache_creation_1h = kwargs.get("cache_creation_tokens_1h", 0)
if cache_creation_5m:
data["cache_creation_input_tokens_5m"] = cache_creation_5m
if cache_creation_1h:
data["cache_creation_input_tokens_1h"] = cache_creation_1h
# 时间指标 # 时间指标
if kwargs.get("response_time_ms") is not None: if kwargs.get("response_time_ms") is not None:
data["response_time_ms"] = kwargs["response_time_ms"] data["response_time_ms"] = kwargs["response_time_ms"]

View File

@@ -142,6 +142,76 @@ class TestDefaultBillingRuleGenerator:
assert result.status == "complete" assert result.status == "complete"
assert abs(float(result.cost) - 0.0005) < 1e-9 assert abs(float(result.cost) - 0.0005) < 1e-9
def test_default_rule_cache_ttl_pricing_overrides_cache_creation_price(self) -> None:
global_model = GlobalModel(
name="ttl-creation-model",
display_name="TTL Creation Model",
is_active=True,
default_price_per_request=0.0,
default_tiered_pricing={
"tiers": [
{
"up_to": None,
"input_price_per_1m": 3.0,
"output_price_per_1m": 15.0,
"cache_creation_price_per_1m": 3.75,
"cache_ttl_pricing": [
{"ttl_minutes": 5, "cache_creation_price_per_1m": 3.75},
{"ttl_minutes": 60, "cache_creation_price_per_1m": 6.0},
],
}
]
},
)
rule = DefaultBillingRuleGenerator.generate_for_model(
global_model=global_model,
model=None,
task_type="chat",
)
engine = FormulaEngine()
# TTL=5: cache_creation_price_per_1m=3.75
# 1_000_000 * 3.75 / 1M = 3.75
result_5m = engine.evaluate(
expression=rule.expression,
variables=rule.variables,
dimensions={
"input_tokens": 0,
"output_tokens": 0,
"cache_creation_tokens": 1_000_000,
"cache_read_tokens": 0,
"cache_ttl_minutes": 5,
"request_count": 1,
"total_input_context": 0,
},
dimension_mappings=rule.dimension_mappings,
strict_mode=True,
)
assert result_5m.status == "complete"
assert abs(float(result_5m.cost) - 3.75) < 1e-9
# TTL=60: cache_creation_price_per_1m=6.0
# 1_000_000 * 6.0 / 1M = 6.0
result_1h = engine.evaluate(
expression=rule.expression,
variables=rule.variables,
dimensions={
"input_tokens": 0,
"output_tokens": 0,
"cache_creation_tokens": 1_000_000,
"cache_read_tokens": 0,
"cache_ttl_minutes": 60,
"request_count": 1,
"total_input_context": 0,
},
dimension_mappings=rule.dimension_mappings,
strict_mode=True,
)
assert result_1h.status == "complete"
assert abs(float(result_1h.cost) - 6.0) < 1e-9
class TestBillingRuleServiceDefaultFallback: class TestBillingRuleServiceDefaultFallback:
def test_find_rule_returns_default_for_chat_when_no_db_rule(self) -> None: def test_find_rule_returns_default_for_chat_when_no_db_rule(self) -> None: