From c28500f42a05aa7f79021ba7164146f086359b89 Mon Sep 17 00:00:00 2001 From: licoy Date: Fri, 7 Aug 2026 16:38:09 +0800 Subject: [PATCH 001/108] fix(avatar): restore local image selection preview for profile and providers Avatar files were saved without an extension so legacy preview rejected them as non-images; provider icons also ignored file/emoji/url kinds on read-back. Add image extensions, magic-byte fallback, optimistic preview, and provider icon codec so selected images display correctly. Closes #140 --- src-tauri/src/commands/files_page.rs | 136 +++++++++++++++--- src/components/layout/Sidebar.tsx | 10 +- src/components/settings/ProviderDetail.tsx | 15 +- src/components/shared/IconEditor.tsx | 37 ++++- .../shared/__tests__/IconEditor.test.tsx | 125 ++++++++++++++++ src/lib/__tests__/providerIconCodec.test.ts | 59 ++++++++ src/lib/providerIconCodec.ts | 54 +++++++ src/lib/providerIcons.tsx | 77 +++++++++- src/pages/RolesPage.tsx | 8 +- 9 files changed, 473 insertions(+), 48 deletions(-) create mode 100644 src/components/shared/__tests__/IconEditor.test.tsx create mode 100644 src/lib/__tests__/providerIconCodec.test.ts create mode 100644 src/lib/providerIconCodec.ts diff --git a/src-tauri/src/commands/files_page.rs b/src-tauri/src/commands/files_page.rs index 2d41a992..dada99e2 100644 --- a/src-tauri/src/commands/files_page.rs +++ b/src-tauri/src/commands/files_page.rs @@ -218,6 +218,34 @@ fn mime_from_extension(path: &str) -> &'static str { } } +/// Map image MIME types to a storage file extension for avatar saves. +/// Returns None for unsupported avatar MIME types. +fn ext_for_image_mime(mime: &str) -> Option<&'static str> { + match mime { + "image/png" => Some("png"), + "image/jpeg" | "image/jpg" => Some("jpg"), + "image/webp" => Some("webp"), + "image/gif" => Some("gif"), + _ => None, + } +} + +/// Sniff common image MIME types from magic bytes. +/// Used when legacy avatar files were saved without an extension. +fn mime_from_image_magic(bytes: &[u8]) -> Option<&'static str> { + if bytes.starts_with(b"\x89PNG\r\n\x1a\n") { + Some("image/png") + } else if bytes.starts_with(&[0xff, 0xd8, 0xff]) { + Some("image/jpeg") + } else if bytes.starts_with(b"GIF87a") || bytes.starts_with(b"GIF89a") { + Some("image/gif") + } else if bytes.len() >= 12 && bytes.starts_with(b"RIFF") && &bytes[8..12] == b"WEBP" { + Some("image/webp") + } else { + None + } +} + const LEGACY_PREVIEW_MAX_BYTES: usize = 1024 * 1024; static LEGACY_PREVIEW_WARNING: std::sync::Once = std::sync::Once::new(); @@ -255,13 +283,6 @@ fn read_attachment_preview_from_root( return Err("Preview image resolves outside the documents images directory".to_string()); } - let mime = mime_from_extension(&canonical_path.to_string_lossy()); - if !matches!( - mime, - "image/png" | "image/jpeg" | "image/webp" | "image/gif" - ) { - return Err("Legacy preview supports PNG, JPEG, WebP, and GIF only".to_string()); - } let size_bytes = canonical_path .metadata() .map_err(|error| format!("Failed to inspect preview image: {error}"))? @@ -274,6 +295,20 @@ fn read_attachment_preview_from_root( let bytes = std::fs::read(&canonical_path) .map_err(|error| format!("Failed to read preview image: {error}"))?; + + let ext_mime = mime_from_extension(&canonical_path.to_string_lossy()); + let mime = if matches!( + ext_mime, + "image/png" | "image/jpeg" | "image/webp" | "image/gif" + ) { + ext_mime + } else { + // Extensionless legacy avatars (pre-fix) still need to preview. + mime_from_image_magic(&bytes).ok_or_else(|| { + "Legacy preview supports PNG, JPEG, WebP, and GIF only".to_string() + })? + }; + aqbot_core::inline_media::validate_image_bytes(mime, &bytes) .map_err(|error| format!("Invalid preview image: {error}"))?; let b64 = base64::engine::general_purpose::STANDARD.encode(&bytes); @@ -308,15 +343,27 @@ pub async fn read_attachment_preview(file_path: String) -> Result Result { use aqbot_core::file_store::FileStore; + // Normalize aliases so extension + byte validation stay in sync. + let mime_type = match mime_type.as_str() { + "image/jpg" => "image/jpeg".to_string(), + other => other.to_string(), + }; + let ext = ext_for_image_mime(&mime_type).ok_or_else(|| { + format!("Unsupported avatar MIME type: {mime_type} (PNG, JPEG, WebP, GIF only)") + })?; let bytes = base64::engine::general_purpose::STANDARD .decode(&data) .map_err(|e| format!("Invalid base64: {e}"))?; + aqbot_core::inline_media::validate_image_bytes(&mime_type, &bytes) + .map_err(|e| format!("Invalid avatar image: {e}"))?; let store = FileStore::new(); let _file_reference_guard = aqbot_core::repo::stored_file::lock_file_references().await; // Avatar paths are not stored_files rows. Give them a unique physical // name so managed attachment GC can never delete an avatar that happens // to have identical bytes and an identical user-supplied file name. - let avatar_name = format!("avatar-{}", aqbot_core::utils::gen_id()); + // Always include an image extension so legacy preview (mime-from-extension) + // and OS tools can identify the file. + let avatar_name = format!("avatar-{}.{}", aqbot_core::utils::gen_id(), ext); let saved = store .save_file(&bytes, &avatar_name, &mime_type) .map_err(|e| format!("Failed to save avatar: {e}"))?; @@ -1088,34 +1135,51 @@ mod tests { // ── save_avatar_file tests ────────────────────────────────────────── - #[tokio::test] - async fn test_save_avatar_file_returns_relative_path() { + fn sample_png_bytes() -> Vec { // 1x1 red PNG pixel - let png_bytes: &[u8] = &[ + vec![ 0x89, 0x50, 0x4E, 0x47, 0x0D, 0x0A, 0x1A, 0x0A, 0x00, 0x00, 0x00, 0x0D, 0x49, 0x48, 0x44, 0x52, 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01, 0x08, 0x02, 0x00, 0x00, 0x00, 0x90, 0x77, 0x53, 0xDE, 0x00, 0x00, 0x00, 0x0C, 0x49, 0x44, 0x41, 0x54, 0x08, 0xD7, 0x63, 0xF8, 0xCF, 0xC0, 0x00, 0x00, 0x00, 0x02, 0x00, 0x01, 0xE2, 0x21, 0xBC, 0x33, 0x00, 0x00, 0x00, 0x00, 0x49, 0x45, 0x4E, 0x44, 0xAE, 0x42, 0x60, 0x82, - ]; - let b64 = base64::engine::general_purpose::STANDARD.encode(png_bytes); + ] + } + + #[test] + fn ext_for_image_mime_maps_supported_types() { + assert_eq!(ext_for_image_mime("image/png"), Some("png")); + assert_eq!(ext_for_image_mime("image/jpeg"), Some("jpg")); + assert_eq!(ext_for_image_mime("image/webp"), Some("webp")); + assert_eq!(ext_for_image_mime("image/gif"), Some("gif")); + assert_eq!(ext_for_image_mime("image/svg+xml"), None); + assert_eq!(ext_for_image_mime("application/octet-stream"), None); + } - // Use a temp dir so we don't pollute the real ~/Documents/aqbot + #[test] + fn test_save_avatar_file_returns_relative_path_with_extension() { + let png_bytes = sample_png_bytes(); let tmp = make_temp_app_data_dir(); - std::fs::create_dir_all(&tmp).unwrap(); + std::fs::create_dir_all(tmp.join("images")).unwrap(); - // Save via FileStore directly (mirrors command logic without the Tauri runtime) + // Mirror command logic (extension from mime + unique avatar name) + let ext = ext_for_image_mime("image/png").expect("png is supported"); + let avatar_name = format!("avatar-test.{}", ext); let store = aqbot_core::file_store::FileStore::with_root(tmp.clone()); - let decoded = base64::engine::general_purpose::STANDARD - .decode(&b64) + let saved = store + .save_file(&png_bytes, &avatar_name, "image/png") .unwrap(); - let saved = store.save_file(&decoded, "avatar", "image/png").unwrap(); assert!( saved.storage_path.starts_with("images/"), "avatar should be stored under images/, got: {}", saved.storage_path ); + assert!( + saved.storage_path.ends_with(".png"), + "avatar path must include image extension for preview, got: {}", + saved.storage_path + ); assert!( saved.storage_path.contains("avatar"), "storage path should contain 'avatar', got: {}", @@ -1123,13 +1187,41 @@ mod tests { ); assert!(tmp.join(&saved.storage_path).exists()); - // Cleanup + // Preview must succeed for the saved path (the #140 regression). + let preview = read_attachment_preview_from_root(&tmp, &saved.storage_path).unwrap(); + assert!( + preview.starts_with("data:image/png;base64,"), + "preview should be a PNG data URI, got: {}", + &preview[..preview.len().min(40)] + ); + let _ = std::fs::remove_dir_all(&tmp); } - #[tokio::test] - async fn test_save_avatar_file_rejects_invalid_base64() { + #[test] + fn legacy_preview_sniffs_extensionless_avatar_bytes() { + let root = make_temp_app_data_dir(); + let images = root.join("images"); + std::fs::create_dir_all(&images).unwrap(); + // Historical bug: avatars were saved without an extension. + std::fs::write(images.join("abcdef_avatar-legacy"), sample_png_bytes()).unwrap(); + + let preview = + read_attachment_preview_from_root(&root, "images/abcdef_avatar-legacy").unwrap(); + assert!(preview.starts_with("data:image/png;base64,")); + + let _ = std::fs::remove_dir_all(root); + } + + #[test] + fn test_save_avatar_file_rejects_invalid_base64() { let result = base64::engine::general_purpose::STANDARD.decode("not-valid-base64!!!"); assert!(result.is_err(), "decoding garbage base64 should fail"); } + + #[test] + fn test_save_avatar_file_rejects_unsupported_mime() { + assert!(ext_for_image_mime("image/bmp").is_none()); + assert!(ext_for_image_mime("text/plain").is_none()); + } } diff --git a/src/components/layout/Sidebar.tsx b/src/components/layout/Sidebar.tsx index e92b60aa..10fe405a 100644 --- a/src/components/layout/Sidebar.tsx +++ b/src/components/layout/Sidebar.tsx @@ -95,7 +95,15 @@ export function Sidebar() { ); } if ((profile.avatarType === 'url' || profile.avatarType === 'file') && profile.avatarValue) { - const src = profile.avatarType === 'file' ? resolvedAvatarSrc : profile.avatarValue; + // Relative paths (images/...) cannot be used as img src; wait for resolved data URI. + const value = profile.avatarValue; + const isDirect = value.slice(0, 64).toLowerCase().startsWith('data:image/') + || value.startsWith('http://') + || value.startsWith('https://') + || value.startsWith('aqbot-media://'); + const src = profile.avatarType === 'file' + ? (resolvedAvatarSrc ?? (isDirect ? value : undefined)) + : value; return ; } return ( diff --git a/src/components/settings/ProviderDetail.tsx b/src/components/settings/ProviderDetail.tsx index ea7894aa..f057be82 100644 --- a/src/components/settings/ProviderDetail.tsx +++ b/src/components/settings/ProviderDetail.tsx @@ -31,6 +31,7 @@ import { useTranslation } from 'react-i18next'; import { invoke } from '@tauri-apps/api/core'; import { useProviderStore, useUIStore } from '@/stores'; import { SmartModelIcon, SmartProviderIcon } from '@/lib/providerIcons'; +import { encodeProviderIcon, parseProviderIcon } from '@/lib/providerIconCodec'; import { getEditableCapabilities, getVisibleModelCapabilities, sanitizeModelCapabilities } from '@/lib/modelCapabilities'; import { IconEditor } from '@/components/shared/IconEditor'; import { DynamicLobeIcon } from '@/components/shared/DynamicLobeIcon'; @@ -1396,20 +1397,14 @@ export function ProviderDetail({ providerId }: ProviderDetailProps) {
{ - if (type === 'model_icon' && value) { - updateProvider(providerId, { icon: value }); - } else if (type === 'emoji' || type === 'url' || type === 'file') { - updateProvider(providerId, { icon: `${type}:${value}` }); - } else { - updateProvider(providerId, { icon: '' }); - } + updateProvider(providerId, { icon: encodeProviderIcon(type, value) }); }} size={40} shape="square" - defaultIcon={} + defaultIcon={} showModelIcons modelIconsDefaultTab="provider" /> diff --git a/src/components/shared/IconEditor.tsx b/src/components/shared/IconEditor.tsx index ca699570..6f404e7c 100644 --- a/src/components/shared/IconEditor.tsx +++ b/src/components/shared/IconEditor.tsx @@ -1,5 +1,5 @@ import { useState, useRef, lazy, Suspense, type ReactNode } from 'react'; -import { Avatar, Dropdown, Input, Button, theme } from 'antd'; +import { Avatar, Dropdown, Input, Button, App, theme } from 'antd'; import type { MenuProps } from 'antd'; import { Smile, Link, FileImage, Trash2, Grid2x2 } from 'lucide-react'; import { useTranslation } from 'react-i18next'; @@ -10,6 +10,19 @@ import { DynamicLobeIcon } from './DynamicLobeIcon'; import { useResolvedAvatarSrc } from '@/hooks/useResolvedAvatarSrc'; import type { AvatarType } from '@/stores/userProfileStore'; +/** Values safe to use as without legacy path resolution. */ +function isDirectImageSource(value: string): boolean { + const prefix = value.slice(0, 64).toLowerCase(); + return ( + prefix.startsWith('data:image/') + || prefix.startsWith('aqbot-media://stored/') + || prefix.startsWith('http://aqbot-media.localhost/stored/') + || prefix.startsWith('https://aqbot-media.localhost/stored/') + || prefix.startsWith('http://') + || prefix.startsWith('https://') + ); +} + const IconPickerModal = lazy(() => import('@/components/settings/IconPickerModal')); export interface IconEditorProps { @@ -55,6 +68,7 @@ export function IconEditor({ modelIconsDefaultTab = 'model', }: IconEditorProps) { const { t } = useTranslation(); + const { message } = App.useApp(); const { token } = theme.useToken(); const fileInputRef = useRef(null); const [showEmojiPicker, setShowEmojiPicker] = useState(false); @@ -70,15 +84,21 @@ export function IconEditor({ reader.onload = async () => { const dataUri = reader.result as string; const match = dataUri.match(/^data:([^;]+);base64,(.+)$/s); + // Optimistic preview: show the data URI immediately so selection feels instant + // and remains visible even if persistence/resolve fails. + onChange('file', dataUri); if (match && isTauri()) { try { - const relativePath = await invoke('save_avatar_file', { data: match[2], mimeType: match[1] }); + const relativePath = await invoke('save_avatar_file', { + data: match[2], + mimeType: match[1], + }); onChange('file', relativePath); - } catch { - onChange('file', dataUri); + } catch (err) { + console.error('save_avatar_file failed:', err); + message.warning(t('userProfile.avatarSaveFailed', '图片已预览,但保存到本地失败,重启后可能丢失')); + // Keep data URI already applied above. } - } else { - onChange('file', dataUri); } }; reader.readAsDataURL(file); @@ -147,7 +167,10 @@ export function IconEditor({ ); } if ((iconType === 'url' || iconType === 'file') && iconValue) { - const src = iconType === 'file' ? (resolvedSrc ?? iconValue) : iconValue; + // Never use a bare relative path (images/...) as — WebView cannot load it. + const src = iconType === 'file' + ? (resolvedSrc ?? (isDirectImageSource(iconValue) ? iconValue : undefined)) + : iconValue; return ; } if (iconType === 'model_icon' && iconValue) { diff --git a/src/components/shared/__tests__/IconEditor.test.tsx b/src/components/shared/__tests__/IconEditor.test.tsx new file mode 100644 index 00000000..b2a52b72 --- /dev/null +++ b/src/components/shared/__tests__/IconEditor.test.tsx @@ -0,0 +1,125 @@ +import { act, fireEvent, render, waitFor } from '@testing-library/react'; +import { App, ConfigProvider } from 'antd'; +import { beforeEach, describe, expect, it, vi } from 'vitest'; + +const mocks = vi.hoisted(() => ({ + invoke: vi.fn(), + tauri: true, + resolvedSrc: undefined as string | undefined, +})); + +vi.mock('@/lib/invoke', () => ({ + invoke: mocks.invoke, + isTauri: () => mocks.tauri, +})); + +vi.mock('@/hooks/useResolvedAvatarSrc', () => ({ + useResolvedAvatarSrc: () => mocks.resolvedSrc, +})); + +vi.mock('@/components/shared/EmojiPicker', () => ({ + EmojiPicker: () => null, +})); + +vi.mock('@/components/shared/DynamicLobeIcon', () => ({ + DynamicLobeIcon: () =>
, +})); + +vi.mock('@/components/settings/IconPickerModal', () => ({ + default: () => null, +})); + +// Avoid pulling emoji-picker-element / heavy icon graphs through relative imports. +vi.mock('../EmojiPicker', () => ({ + EmojiPicker: () => null, +})); + +vi.mock('../DynamicLobeIcon', () => ({ + DynamicLobeIcon: () =>
, +})); + +import { IconEditor } from '../IconEditor'; + +function renderEditor( + props: Partial> & { + onChange?: (type: string | null, value: string | null) => void; + } = {}, +) { + const onChange = props.onChange ?? vi.fn(); + render( + + + + + , + ); + return { onChange }; +} + +describe('IconEditor file selection', () => { + beforeEach(() => { + mocks.invoke.mockReset(); + mocks.tauri = true; + mocks.resolvedSrc = undefined; + }); + + it('optimistically applies data URI then persists relative path', async () => { + mocks.invoke.mockResolvedValueOnce('images/hash_avatar-1.png'); + const onChange = vi.fn(); + renderEditor({ onChange }); + + const input = document.querySelector('input[type="file"]') as HTMLInputElement; + expect(input).toBeTruthy(); + + const file = new File([new Uint8Array([0x89, 0x50, 0x4e, 0x47])], 'a.png', { + type: 'image/png', + }); + const dataUri = 'data:image/png;base64,iVBORw0KGgo='; + const original = FileReader.prototype.readAsDataURL; + FileReader.prototype.readAsDataURL = function (this: FileReader) { + Object.defineProperty(this, 'result', { value: dataUri, configurable: true }); + queueMicrotask(() => this.onload?.({} as ProgressEvent)); + }; + try { + await act(async () => { + fireEvent.change(input, { target: { files: [file] } }); + }); + await waitFor(() => { + expect(onChange).toHaveBeenCalledWith('file', dataUri); + }); + await waitFor(() => { + expect(onChange).toHaveBeenCalledWith('file', 'images/hash_avatar-1.png'); + }); + } finally { + FileReader.prototype.readAsDataURL = original; + } + + expect(mocks.invoke).toHaveBeenCalledWith('save_avatar_file', { + data: 'iVBORw0KGgo=', + mimeType: 'image/png', + }); + }); + + it('does not use relative path as img src before resolve', () => { + renderEditor({ + iconType: 'file', + iconValue: 'images/hash_avatar-1.png', + }); + const img = document.querySelector('img'); + if (img) { + expect(img.getAttribute('src') ?? '').not.toContain('images/hash_avatar'); + } + }); + + it('renders direct data URI without waiting for resolve', () => { + const dataUri = 'data:image/png;base64,iVBORw0KGgo='; + renderEditor({ iconType: 'file', iconValue: dataUri }); + const img = document.querySelector('img'); + expect(img?.getAttribute('src')).toBe(dataUri); + }); +}); diff --git a/src/lib/__tests__/providerIconCodec.test.ts b/src/lib/__tests__/providerIconCodec.test.ts new file mode 100644 index 00000000..40edb7e3 --- /dev/null +++ b/src/lib/__tests__/providerIconCodec.test.ts @@ -0,0 +1,59 @@ +import { describe, expect, it } from 'vitest'; +import { encodeProviderIcon, parseProviderIcon } from '../providerIconCodec'; + +describe('parseProviderIcon', () => { + it('returns null for empty values', () => { + expect(parseProviderIcon(null)).toBeNull(); + expect(parseProviderIcon(undefined)).toBeNull(); + expect(parseProviderIcon('')).toBeNull(); + }); + + it('treats bare keys as model_icon', () => { + expect(parseProviderIcon('OpenAI')).toEqual({ type: 'model_icon', value: 'OpenAI' }); + }); + + it('keeps model:/provider: prefixes as model_icon full value', () => { + expect(parseProviderIcon('model:gpt-4')).toEqual({ type: 'model_icon', value: 'model:gpt-4' }); + expect(parseProviderIcon('provider:OpenAI')).toEqual({ + type: 'model_icon', + value: 'provider:OpenAI', + }); + }); + + it('parses emoji / url / file prefixes', () => { + expect(parseProviderIcon('emoji:😀')).toEqual({ type: 'emoji', value: '😀' }); + expect(parseProviderIcon('url:https://example.com/a.png')).toEqual({ + type: 'url', + value: 'https://example.com/a.png', + }); + expect(parseProviderIcon('file:images/abc_avatar-1.png')).toEqual({ + type: 'file', + value: 'images/abc_avatar-1.png', + }); + }); + + it('preserves data URI after file: prefix', () => { + const data = 'data:image/png;base64,iVBORw0KGgo='; + expect(parseProviderIcon(`file:${data}`)).toEqual({ type: 'file', value: data }); + }); +}); + +describe('encodeProviderIcon', () => { + it('encodes model_icon without prefix rewrite', () => { + expect(encodeProviderIcon('model_icon', 'provider:OpenAI')).toBe('provider:OpenAI'); + expect(encodeProviderIcon('model_icon', 'OpenAI')).toBe('OpenAI'); + }); + + it('prefixes custom kinds', () => { + expect(encodeProviderIcon('emoji', '😀')).toBe('emoji:😀'); + expect(encodeProviderIcon('url', 'https://x.test/a.png')).toBe('url:https://x.test/a.png'); + expect(encodeProviderIcon('file', 'images/a.png')).toBe('file:images/a.png'); + }); + + it('clears on null/empty', () => { + expect(encodeProviderIcon(null, null)).toBe(''); + expect(encodeProviderIcon('file', null)).toBe(''); + expect(encodeProviderIcon('file', '')).toBe(''); + expect(encodeProviderIcon(null, 'x')).toBe(''); + }); +}); diff --git a/src/lib/providerIconCodec.ts b/src/lib/providerIconCodec.ts new file mode 100644 index 00000000..7a96b271 --- /dev/null +++ b/src/lib/providerIconCodec.ts @@ -0,0 +1,54 @@ +/** + * Encode / decode provider.icon values that pack multiple icon kinds into one string. + * + * Storage formats: + * - model/provider lobe icons: `model:OpenAI`, `provider:OpenAI`, or bare `OpenAI` + * - emoji: `emoji:😀` + * - url: `url:https://...` + * - file: `file:images/...` or `file:data:image/png;base64,...` + */ + +export type ProviderIconKind = 'model_icon' | 'emoji' | 'url' | 'file'; + +export interface ParsedProviderIcon { + type: ProviderIconKind; + /** Value passed to IconEditor / renderers (without the type prefix). */ + value: string; +} + +const CUSTOM_PREFIXES = ['emoji', 'url', 'file'] as const; + +/** + * Parse a stored provider.icon string into type + value for IconEditor. + */ +export function parseProviderIcon(icon: string | null | undefined): ParsedProviderIcon | null { + if (!icon) return null; + const sep = icon.indexOf(':'); + if (sep <= 0) { + // Bare lobe icon id / key + return { type: 'model_icon', value: icon }; + } + const prefix = icon.slice(0, sep); + const rest = icon.slice(sep + 1); + if ((CUSTOM_PREFIXES as readonly string[]).includes(prefix) && rest.length > 0) { + return { type: prefix as 'emoji' | 'url' | 'file', value: rest }; + } + // model:xxx / provider:xxx / unknown:xxx → keep full string for DynamicLobeIcon + return { type: 'model_icon', value: icon }; +} + +/** + * Encode IconEditor onChange result into provider.icon storage string. + * Returns empty string when cleared. + */ +export function encodeProviderIcon( + type: string | null, + value: string | null, +): string { + if (!type || value == null || value === '') return ''; + if (type === 'model_icon') return value; + if (type === 'emoji' || type === 'url' || type === 'file') { + return `${type}:${value}`; + } + return ''; +} diff --git a/src/lib/providerIcons.tsx b/src/lib/providerIcons.tsx index ec3235b9..b146a20a 100644 --- a/src/lib/providerIcons.tsx +++ b/src/lib/providerIcons.tsx @@ -1,7 +1,10 @@ import { memo } from 'react'; +import { Avatar } from 'antd'; import type { ProviderConfig } from '@/types'; import { ProviderIcon, ModelIcon, providerMappings, modelMappings } from '@lobehub/icons'; import { DynamicLobeIcon } from '@/components/shared/DynamicLobeIcon'; +import { useResolvedAvatarSrc } from '@/hooks/useResolvedAvatarSrc'; +import { parseProviderIcon } from '@/lib/providerIconCodec'; const SHUAI_API_LOGO_URL = 'https://api.shuaiapi.com/images/logo.svg'; const GPTNB_LOGO_URL = 'https://pic.scdn.app/images/2023/06/26/favicon.png'; @@ -118,8 +121,34 @@ export function getProviderIconKey(provider: ProviderConfig): string { return result.key; } +function ProviderFileIcon({ + value, + size, + shape, +}: { + value: string; + size: number; + shape?: 'circle' | 'square'; +}) { + const resolvedSrc = useResolvedAvatarSrc('file', value); + const direct = + value.slice(0, 64).toLowerCase().startsWith('data:image/') + || value.startsWith('aqbot-media://') + || value.includes('aqbot-media.localhost'); + const src = resolvedSrc ?? (direct ? value : undefined); + return ( + + ); +} + /** - * Two-tier icon component: tries ProviderIcon first, then ModelIcon, then fallback. + * Two-tier icon component: tries custom icon, then ProviderIcon/ModelIcon fallback. + * Custom provider.icon may be emoji/url/file (prefixed) or a lobe model/provider key. */ export const SmartProviderIcon = memo(function SmartProviderIcon({ provider, @@ -132,12 +161,46 @@ export const SmartProviderIcon = memo(function SmartProviderIcon({ type?: 'avatar' | 'color' | 'mono'; shape?: 'circle' | 'square'; }) { - if (provider.icon) { - const [, key] = provider.icon.includes(':') - ? (provider.icon.split(':', 2) as [string, string]) - : ['model' as const, provider.icon]; - // key is a toc `id` (e.g., "Ai302", "OpenAI") — use DynamicLobeIcon for reliable rendering - return ; + const parsed = parseProviderIcon(provider.icon); + if (parsed) { + if (parsed.type === 'emoji') { + const borderRadius = shape === 'square' ? Math.floor(size * 0.1) : '50%'; + return ( +
+ {parsed.value} +
+ ); + } + if (parsed.type === 'url') { + return ( + + ); + } + if (parsed.type === 'file') { + return ; + } + // model_icon: value is `group:id` or bare id + const iconId = parsed.value.includes(':') + ? parsed.value.slice(parsed.value.indexOf(':') + 1) + : parsed.value; + return ; } const builtinLogoUrl = provider.builtin_id ? BUILTIN_LOGO_URLS[provider.builtin_id] diff --git a/src/pages/RolesPage.tsx b/src/pages/RolesPage.tsx index 9361f45d..07d7fe8a 100644 --- a/src/pages/RolesPage.tsx +++ b/src/pages/RolesPage.tsx @@ -190,7 +190,13 @@ function RoleAvatar({ role }: { role: Pick; } return ( From 55b3005cf565d2430aebba6dd6dfd5bad2cef87c Mon Sep 17 00:00:00 2001 From: licoy Date: Fri, 7 Aug 2026 17:21:47 +0800 Subject: [PATCH 002/108] =?UTF-8?q?feat(chat):=20=E4=BC=98=E5=8C=96?= =?UTF-8?q?=E6=B6=88=E6=81=AF=E5=88=86=E4=BA=AB=E5=AF=BC=E5=87=BA=E4=BD=93?= =?UTF-8?q?=E9=AA=8C=E4=B8=8E=E5=91=88=E7=8E=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 说明: - PNG 改为数据驱动离屏渲染,Markdown 正确排版,避免视口裁切与操作栏噪点 - 导出展示统一 i18n 角色文案;助手标题用模型名;消息底部仅显示时间 - 文本/Markdown/JSON 与 PNG 对齐说话人与时间;文件名追加 ` - {timestamp}` 避免覆盖 - 分享选择支持点击消息选中,主按钮一键导出 PNG,并抽出共享文件名清洗逻辑 Closes #139 --- src/components/chat/ChatSidebar.tsx | 41 +- src/components/chat/ChatView.tsx | 177 +++++-- .../feedback.phase-c.output-controls.test.tsx | 6 + src/i18n/locales/en-US.json | 3 + src/i18n/locales/zh-CN.json | 3 + src/lib/__tests__/exportChat.test.ts | 189 +++++++ .../__tests__/exportChatPresentation.test.ts | 81 +++ src/lib/__tests__/filename.test.ts | 44 ++ src/lib/chatImageActions.ts | 9 +- src/lib/exportChat.ts | 497 +++++++++++++++++- src/lib/exportChatPresentation.ts | 79 +++ src/lib/filename.ts | 45 ++ 12 files changed, 1096 insertions(+), 78 deletions(-) create mode 100644 src/lib/__tests__/exportChat.test.ts create mode 100644 src/lib/__tests__/exportChatPresentation.test.ts create mode 100644 src/lib/__tests__/filename.test.ts create mode 100644 src/lib/exportChatPresentation.ts create mode 100644 src/lib/filename.ts diff --git a/src/components/chat/ChatSidebar.tsx b/src/components/chat/ChatSidebar.tsx index aeaa0561..eef5bc61 100644 --- a/src/components/chat/ChatSidebar.tsx +++ b/src/components/chat/ChatSidebar.tsx @@ -2,14 +2,16 @@ import { useState, useMemo, useCallback, useEffect, useRef, memo } from 'react' import { Button, Input, App, theme, Tooltip, Checkbox, Dropdown, Empty } from 'antd' import { MessageSquarePlus, Search, Archive, ListTodo, Trash2, Pencil, Share, Pin, PinOff, Loader, X, Undo2, ArrowLeft, FileImage, FileCode, FileType, FileText, FolderPlus, FolderOpen, GripVertical, ChevronRight, MessageSquareText, Sparkles } from 'lucide-react' import { exportAsMarkdown, exportAsText, exportMessagesAsPNG, exportAsJSON } from '@/lib/exportChat' +import { buildExportOptions } from '@/lib/exportChatPresentation' import { invoke } from '@/lib/invoke' +import { useUserProfileStore } from '@/stores/userProfileStore' +import { useResolvedAvatarSrc } from '@/hooks/useResolvedAvatarSrc' import type { ConversationItemType } from '@ant-design/x/es/conversations/interface' import { useTranslation } from 'react-i18next' import { useConversationStore, useProviderStore, useSettingsStore, useCategoryStore } from '@/stores' import { getShortcutBinding, formatShortcutForDisplay } from '@/lib/shortcuts' import type { ShortcutAction } from '@/lib/shortcuts' import type { Conversation, Message, ConversationCategory } from '@/types' -import { useResolvedAvatarSrc } from '@/hooks/useResolvedAvatarSrc' import type { AvatarType } from '@/stores/userProfileStore' import { CategoryEditModal, type CategoryEditFormData } from './CategoryEditModal' import { ConversationIcon } from './ConversationIcon' @@ -173,6 +175,7 @@ export function ChatSidebar() { const providers = useProviderStore((s) => s.providers) const settings = useSettingsStore((s) => s.settings) const settingsLoading = useSettingsStore((s) => s.loading) + const profile = useUserProfileStore((s) => s.profile) const categories = useCategoryStore((s) => s.categories) const ensureCategoriesLoaded = useCategoryStore((s) => s.ensureCategoriesLoaded) @@ -935,6 +938,19 @@ export function ChatSidebar() { const buildExportChildren = useCallback( (convId: string, title: string) => { + const conv = conversationById.get(convId) + const exportOptions = buildExportOptions({ + userName: profile.name, + theme: { + colorPrimary: token.colorPrimary, + colorPrimaryBg: token.colorPrimaryBg, + colorPrimaryBorder: token.colorPrimaryBorder, + colorFillSecondary: token.colorFillSecondary, + }, + providers, + conversationModelId: conv?.model_id, + conversationProviderId: conv?.provider_id, + }) return [ { key: 'export-png', @@ -945,7 +961,10 @@ export function ChatSidebar() { const msgs = await invoke('list_messages', { conversationId: convId }) const shareable = msgs.filter((m) => m.role === 'user' || m.role === 'assistant') if (shareable.length === 0) { messageApi.warning(t('chat.noMessages')); return } - const ok = await exportMessagesAsPNG(shareable, title, { includeThinking: false }) + const ok = await exportMessagesAsPNG(shareable, title, { + ...exportOptions, + includeThinking: false, + }) if (ok) messageApi.success(t('chat.exportSuccess')) } catch (e) { console.error('Export PNG failed:', e) @@ -961,7 +980,7 @@ export function ChatSidebar() { try { const msgs = await invoke('list_messages', { conversationId: convId }) if (msgs.length === 0) { messageApi.warning(t('chat.noMessages')); return } - const ok = await exportAsMarkdown(msgs, title) + const ok = await exportAsMarkdown(msgs, title, exportOptions) if (ok) messageApi.success(t('chat.exportSuccess')) } catch (e) { console.error('Export MD failed:', e) @@ -977,7 +996,7 @@ export function ChatSidebar() { try { const msgs = await invoke('list_messages', { conversationId: convId }) if (msgs.length === 0) { messageApi.warning(t('chat.noMessages')); return } - const ok = await exportAsText(msgs, title) + const ok = await exportAsText(msgs, title, exportOptions) if (ok) messageApi.success(t('chat.exportSuccess')) } catch (e) { console.error('Export TXT failed:', e) @@ -993,7 +1012,7 @@ export function ChatSidebar() { try { const msgs = await invoke('list_messages', { conversationId: convId }) if (msgs.length === 0) { messageApi.warning(t('chat.noMessages')); return } - const ok = await exportAsJSON(msgs, title) + const ok = await exportAsJSON(msgs, title, exportOptions) if (ok) messageApi.success(t('chat.exportSuccess')) } catch (e) { console.error('Export JSON failed:', e) @@ -1003,7 +1022,17 @@ export function ChatSidebar() { }, ] }, - [messageApi, t], + [ + conversationById, + messageApi, + profile.name, + providers, + t, + token.colorFillSecondary, + token.colorPrimary, + token.colorPrimaryBg, + token.colorPrimaryBorder, + ], ) const menuConfig = useCallback( diff --git a/src/components/chat/ChatView.tsx b/src/components/chat/ChatView.tsx index 1a1cc4a8..f67e7393 100644 --- a/src/components/chat/ChatView.tsx +++ b/src/components/chat/ChatView.tsx @@ -1806,6 +1806,7 @@ function AssistantFooter({ // ── Export helpers ────────────────────────────────────────────────────── import { copyTranscript, exportMessagesAsPNG, exportAsMarkdown, exportAsJSON, exportAsText } from '@/lib/exportChat'; +import { buildExportOptions } from '@/lib/exportChatPresentation'; // ── Stats Popover ────────────────────────────────────────────────────── @@ -2721,6 +2722,44 @@ export function ChatView() { )); }, []); + /** Click message body/header (not interactive controls) to toggle share selection. */ + const handleShareSelectableClick = useCallback((messageId: string | undefined, e: React.MouseEvent) => { + if (!shareSelectMode || !messageId) return; + const target = e.target as HTMLElement | null; + if (target?.closest( + 'a, button, input, textarea, .ant-checkbox-wrapper, .ant-checkbox, [data-share-ignore="true"]', + )) { + return; + } + toggleShareMessage(messageId); + }, [shareSelectMode, toggleShareMessage]); + + const wrapShareSelectableContent = useCallback((messageId: string | undefined, node: React.ReactNode) => { + if (!shareSelectMode || !messageId) return node; + return ( +
handleShareSelectableClick(messageId, e)} + style={{ cursor: 'pointer' }} + > + {node} +
+ ); + }, [handleShareSelectableClick, shareSelectMode]); + + const getShareSelectBubbleStyles = useCallback((messageId: string | undefined) => { + if (!shareSelectMode || !messageId) return undefined; + const selected = selectedShareMessageIds.includes(messageId); + return { + root: { cursor: 'pointer' as const }, + content: { + cursor: 'pointer' as const, + boxShadow: selected ? `0 0 0 2px ${token.colorPrimary}` : undefined, + transition: 'box-shadow 0.15s ease', + }, + }; + }, [selectedShareMessageIds, shareSelectMode, token.colorPrimary]); + const selectAllShareMessages = useCallback(() => { setSelectedShareMessageIds(shareableMessages.map((m) => m.id)); }, [shareableMessages]); @@ -2730,6 +2769,31 @@ export function ChatView() { return shareableMessages.filter((m) => selected.has(m.id)); }, [selectedShareMessageIds, shareableMessages]); + const buildChatExportOptions = useCallback((includeThinking = false) => ({ + ...buildExportOptions({ + userName: profile.name, + theme: { + colorPrimary: token.colorPrimary, + colorPrimaryBg: token.colorPrimaryBg, + colorPrimaryBorder: token.colorPrimaryBorder, + colorFillSecondary: token.colorFillSecondary, + }, + providers, + conversationModelId: activeConversation?.model_id, + conversationProviderId: activeConversation?.provider_id, + }), + includeThinking, + }), [ + activeConversation?.model_id, + activeConversation?.provider_id, + profile.name, + providers, + token.colorFillSecondary, + token.colorPrimary, + token.colorPrimaryBg, + token.colorPrimaryBorder, + ]); + const exportSelectedShare = useCallback(async (format: 'png' | 'md' | 'copy-md') => { const selected = getSelectedShareMessagesOrdered(); if (selected.length === 0) { @@ -2737,22 +2801,23 @@ export function ChatView() { return; } const title = activeConversation?.title ?? 'chat'; + const exportOptions = buildChatExportOptions(false); setShareExporting(true); try { if (format === 'png') { - const ok = await exportMessagesAsPNG(selected, title, { includeThinking: false }); + const ok = await exportMessagesAsPNG(selected, title, exportOptions); if (ok) { messageApi.success(t('chat.exportSuccess')); exitShareSelectMode(); } } else if (format === 'md') { - const ok = await exportAsMarkdown(selected, title, { includeThinking: false }); + const ok = await exportAsMarkdown(selected, title, exportOptions); if (ok) { messageApi.success(t('chat.exportSuccess')); exitShareSelectMode(); } } else { - const ok = await copyTranscript(selected, title, 'markdown', { includeThinking: false }); + const ok = await copyTranscript(selected, title, 'markdown', exportOptions); if (ok) { messageApi.success(t('chat.copied')); exitShareSelectMode(); @@ -2764,7 +2829,7 @@ export function ChatView() { } finally { setShareExporting(false); } - }, [activeConversation?.title, exitShareSelectMode, getSelectedShareMessagesOrdered, messageApi, t]); + }, [activeConversation?.title, buildChatExportOptions, exitShareSelectMode, getSelectedShareMessagesOrdered, messageApi, t]); const exportMenuItems = useMemo( () => [ @@ -2788,7 +2853,12 @@ export function ChatView() { try { const transcript = await loadCompleteTranscript(); if (transcript.length === 0) { messageApi.warning(t('chat.noMessages')); return; } - const ok = await copyTranscript(transcript, activeConversation?.title ?? 'chat', 'markdown', { includeThinking: false }); + const ok = await copyTranscript( + transcript, + activeConversation?.title ?? 'chat', + 'markdown', + buildChatExportOptions(false), + ); if (ok) messageApi.success(t('chat.copied')); } catch (e) { console.error('Copy MD failed:', e); messageApi.error(t('chat.exportFailed')); } }, @@ -2803,7 +2873,11 @@ export function ChatView() { const transcript = await loadCompleteTranscript(); const shareable = transcript.filter((m) => m.role === 'user' || m.role === 'assistant'); if (shareable.length === 0) { messageApi.warning(t('chat.noMessages')); return; } - const ok = await exportMessagesAsPNG(shareable, activeConversation?.title ?? 'chat', { includeThinking: false }); + const ok = await exportMessagesAsPNG( + shareable, + activeConversation?.title ?? 'chat', + buildChatExportOptions(false), + ); if (ok) messageApi.success(t('chat.exportSuccess')); } catch (e) { console.error('Export PNG failed:', e); messageApi.error(t('chat.exportFailed')); } }, @@ -2816,7 +2890,11 @@ export function ChatView() { try { const transcript = await loadCompleteTranscript(); if (transcript.length === 0) { messageApi.warning(t('chat.noMessages')); return; } - const ok = await exportAsMarkdown(transcript, activeConversation?.title ?? 'chat'); + const ok = await exportAsMarkdown( + transcript, + activeConversation?.title ?? 'chat', + buildChatExportOptions(true), + ); if (ok) messageApi.success(t('chat.exportSuccess')); } catch (e) { console.error('Export MD failed:', e); messageApi.error(t('chat.exportFailed')); } }, @@ -2829,7 +2907,11 @@ export function ChatView() { try { const transcript = await loadCompleteTranscript(); if (transcript.length === 0) { messageApi.warning(t('chat.noMessages')); return; } - const ok = await exportAsMarkdown(transcript, activeConversation?.title ?? 'chat', { includeThinking: false }); + const ok = await exportAsMarkdown( + transcript, + activeConversation?.title ?? 'chat', + buildChatExportOptions(false), + ); if (ok) messageApi.success(t('chat.exportSuccess')); } catch (e) { console.error('Export MD (no thinking) failed:', e); messageApi.error(t('chat.exportFailed')); } }, @@ -2842,7 +2924,11 @@ export function ChatView() { try { const transcript = await loadCompleteTranscript(); if (transcript.length === 0) { messageApi.warning(t('chat.noMessages')); return; } - const ok = await exportAsText(transcript, activeConversation?.title ?? 'chat'); + const ok = await exportAsText( + transcript, + activeConversation?.title ?? 'chat', + buildChatExportOptions(true), + ); if (ok) messageApi.success(t('chat.exportSuccess')); } catch (e) { console.error('Export TXT failed:', e); messageApi.error(t('chat.exportFailed')); } }, @@ -2855,7 +2941,11 @@ export function ChatView() { try { const transcript = await loadCompleteTranscript(); if (transcript.length === 0) { messageApi.warning(t('chat.noMessages')); return; } - const ok = await exportAsText(transcript, activeConversation?.title ?? 'chat', { includeThinking: false }); + const ok = await exportAsText( + transcript, + activeConversation?.title ?? 'chat', + buildChatExportOptions(false), + ); if (ok) messageApi.success(t('chat.exportSuccess')); } catch (e) { console.error('Export TXT (no thinking) failed:', e); messageApi.error(t('chat.exportFailed')); } }, @@ -2868,7 +2958,11 @@ export function ChatView() { try { const transcript = await loadCompleteTranscript(); if (transcript.length === 0) { messageApi.warning(t('chat.noMessages')); return; } - const ok = await exportAsJSON(transcript, activeConversation?.title ?? 'chat'); + const ok = await exportAsJSON( + transcript, + activeConversation?.title ?? 'chat', + buildChatExportOptions(true), + ); if (ok) messageApi.success(t('chat.exportSuccess')); } catch (e) { console.error('Export JSON failed:', e); messageApi.error(t('chat.exportFailed')); } }, @@ -2881,13 +2975,17 @@ export function ChatView() { try { const transcript = await loadCompleteTranscript(); if (transcript.length === 0) { messageApi.warning(t('chat.noMessages')); return; } - const ok = await exportAsJSON(transcript, activeConversation?.title ?? 'chat', { includeThinking: false }); + const ok = await exportAsJSON( + transcript, + activeConversation?.title ?? 'chat', + buildChatExportOptions(false), + ); if (ok) messageApi.success(t('chat.exportSuccess')); } catch (e) { console.error('Export JSON (no thinking) failed:', e); messageApi.error(t('chat.exportFailed')); } }, }, ], - [activeConversation, enterShareSelectMode, loadCompleteTranscript, messageApi, shareableMessages.length, t], + [activeConversation, buildChatExportOptions, enterShareSelectMode, loadCompleteTranscript, messageApi, shareableMessages.length, t], ); // ── Welcome prompt items ─────────────────────────────────────────── @@ -3650,8 +3748,9 @@ export function ChatView() { placement: 'end' as const, ...getBubbleVariant(true), avatar: userAvatar, + styles: getShareSelectBubbleStyles(msg?.id), contentRender: attachments.length > 0 - ? (content: string) => ( + ? (content: string) => wrapShareSelectableContent(msg?.id, (
{renderUserContent(content, 'right')} @@ -3665,16 +3764,16 @@ export function ChatView() { ))}
- ) - : (content: string) => ( + )) + : (content: string) => wrapShareSelectableContent(msg?.id, ( <> {renderUserContent(content)} - ), + )), header: ( -
-
+
handleShareSelectableClick(msg?.id, e)}> +
{shareSelectMode && msg && ( ), }; - }, [activeConversationId, codeBlockDarkTheme, codeBlockLightTheme, codeBlockThemes, deleteMessageGroup, formatTime, getBubbleVariant, handleEditMessage, isDarkMode, messageApi, messageById, profile.name, regenerateMessage, selectedShareMessageIds, settings.code_font_family, settings.render_user_markdown, shareSelectMode, t, toggleShareMessage, token.colorError, token.colorPrimary, userAvatar]); + }, [activeConversationId, codeBlockDarkTheme, codeBlockLightTheme, codeBlockThemes, deleteMessageGroup, formatTime, getBubbleVariant, getShareSelectBubbleStyles, handleEditMessage, handleShareSelectableClick, isDarkMode, messageApi, messageById, profile.name, regenerateMessage, selectedShareMessageIds, settings.code_font_family, settings.render_user_markdown, shareSelectMode, t, toggleShareMessage, token.colorError, token.colorPrimary, userAvatar, wrapShareSelectableContent]); const renderStreamingStatusIndicator = useCallback(( activity: StreamActivity | undefined, @@ -3892,6 +3991,7 @@ export function ChatView() { ...getBubbleVariant(false), avatar: isNonTabsMultiModel ? undefined : renderConvIconForChat(32, msg?.model_id), loading: bubbleLoading, + styles: getShareSelectBubbleStyles(msg?.id), contentRender: (content: string) => { const baseRenderContent = typeof content === 'string' && content.length > 0 ? content @@ -4076,7 +4176,7 @@ export function ChatView() { }; if (isStreaming && msg?.id) { - return ( + return wrapShareSelectableContent(msg.id, ( - ); + )); } - return renderContentNode(baseRenderContent); + return wrapShareSelectableContent(msg?.id, renderContentNode(baseRenderContent)); }, header: (() => { if (isNonTabsMultiModel && !shareSelectMode) return null; const { modelName, providerName } = getModelDisplayInfo(msg?.model_id, msg?.provider_id); return ( -
+
handleShareSelectableClick(msg?.id, e)} + >
{shareSelectMode && msg && ( ) : null, }; - }, [activeConversation, activeConversationId, activeMessages, agentPendingPermissions, agentToolCalls, aiContentNodesById, assistantByParentId, codeBlockDarkTheme, codeBlockLightTheme, codeBlockThemes, currentMessageVersionsByParentId, deleteMessage, displayModeOverrides, displayVersionOverrides, formatTime, getBubbleVariant, getModelDisplayInfo, handleBranchDisplayedVersion, handleDisplayModeOverride, handleDisplayVersionOverride, handleEditMessage, handleGeneratedVersionCreated, handleMultiModelDetected, handleRegenerateDisplayedVersion, handleSetContextVersion, handleSwitchDisplayedVersionModel, isDarkMode, messageById, messages, multiModelDoneMessageIds, multiModelParentId, multiModelResponseParents, ragDisplayByMessageId, renderConvIconForChat, renderStreamingStatusIndicator, searchDisplayByMessageId, selectedShareMessageIds, settings, shareSelectMode, streamActivityByMessageId, streaming, streamingMessageId, switchMessageVersion, t, toggleShareMessage, token.colorPrimary, token.colorTextDescription]); + }, [activeConversation, activeConversationId, activeMessages, agentPendingPermissions, agentToolCalls, aiContentNodesById, assistantByParentId, codeBlockDarkTheme, codeBlockLightTheme, codeBlockThemes, currentMessageVersionsByParentId, deleteMessage, displayModeOverrides, displayVersionOverrides, formatTime, getBubbleVariant, getModelDisplayInfo, getShareSelectBubbleStyles, handleBranchDisplayedVersion, handleDisplayModeOverride, handleDisplayVersionOverride, handleEditMessage, handleGeneratedVersionCreated, handleMultiModelDetected, handleRegenerateDisplayedVersion, handleSetContextVersion, handleShareSelectableClick, handleSwitchDisplayedVersionModel, isDarkMode, messageById, messages, multiModelDoneMessageIds, multiModelParentId, multiModelResponseParents, ragDisplayByMessageId, renderConvIconForChat, renderStreamingStatusIndicator, searchDisplayByMessageId, selectedShareMessageIds, settings, shareSelectMode, streamActivityByMessageId, streaming, streamingMessageId, switchMessageVersion, t, toggleShareMessage, token.colorPrimary, token.colorTextDescription, wrapShareSelectableContent]); const contextClearRole = useCallback((bubbleData: BubbleItemType) => { const msgId = String(bubbleData.content ?? ''); @@ -4642,16 +4745,19 @@ export function ChatView() { + , - label: t('chat.exportPng'), - disabled: selectedShareMessageIds.length === 0 || shareExporting, - onClick: () => { void exportSelectedShare('png'); }, - }, { key: 'md', icon: , @@ -4674,13 +4780,10 @@ export function ChatView() { > diff --git a/src/components/chat/__tests__/feedback.phase-c.output-controls.test.tsx b/src/components/chat/__tests__/feedback.phase-c.output-controls.test.tsx index db9a494b..d214cd54 100644 --- a/src/components/chat/__tests__/feedback.phase-c.output-controls.test.tsx +++ b/src/components/chat/__tests__/feedback.phase-c.output-controls.test.tsx @@ -39,8 +39,14 @@ describe('Phase C output control regressions', () => { expect(chatView).toContain("key: 'select-share'"); expect(chatView).toContain('exportMessagesAsPNG'); expect(chatView).toContain('shareSelectMode'); + expect(chatView).toContain('handleShareSelectableClick'); + expect(chatView).toContain('wrapShareSelectableContent'); + // Primary CTA exports PNG directly (not only via dropdown) + expect(chatView).toContain("void exportSelectedShare('png')"); expect(exportChat).toContain('exportMessagesAsPNG'); expect(exportChat).toContain('prepareClonedExportRoot'); + expect(exportChat).toContain('renderExportMarkdownHtml'); + expect(exportChat).toContain('sanitizeExportFilename'); }); it('lets export helpers optionally strip thinking content before saving or copying', () => { diff --git a/src/i18n/locales/en-US.json b/src/i18n/locales/en-US.json index 49a7baf3..b8032ada 100644 --- a/src/i18n/locales/en-US.json +++ b/src/i18n/locales/en-US.json @@ -143,6 +143,8 @@ "stop": "Stop", "commandHint": "Type / for commands", "you": "You", + "assistant": "Assistant", + "system": "System", "archive": "Archive", "archived": "Archived", "unarchive": "Unarchive", @@ -225,6 +227,7 @@ "shareSelectNone": "Select messages to share first", "shareSelectedCount": "{{count}} selected", "shareSelectAll": "Select all", + "shareMoreFormats": "More", "mcp": { "title": "MCP Tools", "noServers": "Please configure and enable MCP servers in settings first", diff --git a/src/i18n/locales/zh-CN.json b/src/i18n/locales/zh-CN.json index e7471a91..e9929802 100644 --- a/src/i18n/locales/zh-CN.json +++ b/src/i18n/locales/zh-CN.json @@ -143,6 +143,8 @@ "stop": "停止", "commandHint": "输入 / 使用命令", "you": "你", + "assistant": "助手", + "system": "系统", "archive": "归档", "archived": "已归档", "unarchive": "取消归档", @@ -225,6 +227,7 @@ "shareSelectNone": "请先选择要分享的消息", "shareSelectedCount": "已选 {{count}} 条", "shareSelectAll": "全选", + "shareMoreFormats": "更多", "mcp": { "title": "MCP 工具", "noServers": "请先在设置中配置并启用 MCP 服务器", diff --git a/src/lib/__tests__/exportChat.test.ts b/src/lib/__tests__/exportChat.test.ts new file mode 100644 index 00000000..4023c8e7 --- /dev/null +++ b/src/lib/__tests__/exportChat.test.ts @@ -0,0 +1,189 @@ +import { describe, expect, it } from 'vitest'; +import { + renderExportMarkdownHtml, + buildMarkdownTranscript, + buildTextTranscript, + buildJsonTranscript, + resolveExportSpeakerLabel, +} from '../exportChat'; +import type { Message } from '@/types'; + +function makeMessage(partial: Partial & Pick): Message { + return { + conversation_id: 'c1', + status: 'complete', + created_at: 1_704_067_200_000, // 2024-01-01 UTC-ish; formatTime will map it + updated_at: 1_704_067_200_000, + ...partial, + } as Message; +} + +describe('renderExportMarkdownHtml', () => { + it('renders headings, emphasis, and lists instead of raw markdown source', () => { + const html = renderExportMarkdownHtml([ + '# Title', + '', + 'Hello **world** and `code`', + '', + '- item one', + '- item two', + ].join('\n')); + + expect(html).toContain('world'); + expect(html).toContain('export-code-inline'); + expect(html).toContain(' { + const html = renderExportMarkdownHtml('```ts\nconst x = 1 < 2\n```'); + expect(html).toContain('export-code-block'); + expect(html).toContain('const x = 1 < 2'); + expect(html).not.toContain('```'); + }); + + it('does not pass through raw script/html tags', () => { + const html = renderExportMarkdownHtml('Hello **ok**'); + // Parser strips or escapes untrusted HTML; never emit executable tags. + expect(html).not.toMatch(/ok'); + }); + + it('escapes HTML special characters in code spans', () => { + const html = renderExportMarkdownHtml('use `a < b && c > d`'); + expect(html).toContain('export-code-inline'); + expect(html).toContain('<'); + expect(html).toContain('&&'); + expect(html).toContain('>'); + }); + + it('returns empty string for blank input', () => { + expect(renderExportMarkdownHtml(' ')).toBe(''); + }); +}); + +describe('resolveExportSpeakerLabel', () => { + it('uses userName for user and model label for assistant', () => { + const opts = { + roleLabels: { user: '你', assistant: '助手', system: '系统' }, + userName: '小明', + getModelLabel: () => 'GPT-4o', + }; + expect(resolveExportSpeakerLabel(makeMessage({ id: 'u', role: 'user', content: 'hi' }), opts)).toBe('小明'); + expect(resolveExportSpeakerLabel(makeMessage({ id: 'a', role: 'assistant', content: 'yo' }), opts)).toBe('GPT-4o'); + expect(resolveExportSpeakerLabel(makeMessage({ id: 's', role: 'system', content: 'sys' }), opts)).toBe('系统'); + }); + + it('falls back to assistant label when model unknown', () => { + const opts = { + roleLabels: { user: 'You', assistant: 'Assistant', system: 'System' }, + }; + expect(resolveExportSpeakerLabel(makeMessage({ id: 'a', role: 'assistant', content: 'yo' }), opts)).toBe('Assistant'); + }); +}); + +describe('buildMarkdownTranscript', () => { + it('optionally strips thinking via includeThinking false', () => { + const messages = [ + makeMessage({ + id: 'm1', + role: 'assistant', + content: 'secret\n\nVisible answer', + }), + ]; + const withThink = buildMarkdownTranscript(messages, 't'); + const withoutThink = buildMarkdownTranscript(messages, 't', { includeThinking: false }); + expect(withThink).toContain('secret'); + expect(withoutThink).not.toContain('secret'); + expect(withoutThink).toContain('Visible answer'); + }); + + it('uses i18n speaker labels, model name, and time footer', () => { + const messages = [ + makeMessage({ id: 'u1', role: 'user', content: '你好', created_at: 1_704_067_200_000 }), + makeMessage({ + id: 'a1', + role: 'assistant', + content: '世界', + model_id: 'gpt-4o', + created_at: 1_704_067_260_000, + }), + ]; + const md = buildMarkdownTranscript(messages, '会话', { + roleLabels: { user: '你', assistant: '助手', system: '系统' }, + userName: '小明', + getModelLabel: (m) => (m.role === 'assistant' ? 'GPT-4o' : undefined), + formatTime: () => '12:00:00', + }); + + expect(md).toContain('# 会话'); + expect(md).toContain('## 小明'); + expect(md).toContain('## GPT-4o'); + expect(md).not.toContain('## 助手'); + expect(md).not.toContain('## Assistant'); + expect(md).toContain('你好'); + expect(md).toContain('世界'); + // time appears after each message body + expect(md.match(/12:00:00/g)?.length).toBe(2); + }); +}); + +describe('buildTextTranscript', () => { + it('mirrors markdown speaker/time presentation', () => { + const messages = [ + makeMessage({ id: 'u1', role: 'user', content: 'hi' }), + makeMessage({ id: 'a1', role: 'assistant', content: 'hey' }), + ]; + const text = buildTextTranscript(messages, 'Chat', { + roleLabels: { user: '你', assistant: '助手', system: '系统' }, + userName: 'Alice', + getModelLabel: () => 'Claude', + formatTime: () => '09:30:00', + }); + + expect(text).toContain('[Alice]'); + expect(text).toContain('[Claude]'); + expect(text).not.toContain('[助手]'); + expect(text).not.toContain('[Assistant]'); + expect(text).toContain('09:30:00'); + }); +}); + +describe('buildJsonTranscript', () => { + it('includes display name (model for assistant) and formatted time', () => { + const messages = [ + makeMessage({ id: 'u1', role: 'user', content: '你好' }), + makeMessage({ id: 'a1', role: 'assistant', content: '世界', model_id: 'gpt-4o' }), + ]; + const json = JSON.parse(buildJsonTranscript(messages, '会话', { + roleLabels: { user: '你', assistant: '助手', system: '系统' }, + userName: '小明', + getModelLabel: (m) => (m.role === 'assistant' ? 'GPT-4o' : undefined), + formatTime: () => '12:00:00', + includeThinking: false, + })); + + expect(json.title).toBe('会话'); + expect(json.messages).toHaveLength(2); + expect(json.messages[0]).toMatchObject({ + role: 'user', + name: '小明', + content: '你好', + time: '12:00:00', + }); + expect(json.messages[1]).toMatchObject({ + role: 'assistant', + name: 'GPT-4o', + content: '世界', + time: '12:00:00', + }); + expect(json.messages[1].name).not.toBe('助手'); + expect(json.messages[1].thinking).toBeUndefined(); + }); +}); diff --git a/src/lib/__tests__/exportChatPresentation.test.ts b/src/lib/__tests__/exportChatPresentation.test.ts new file mode 100644 index 00000000..f15ef20a --- /dev/null +++ b/src/lib/__tests__/exportChatPresentation.test.ts @@ -0,0 +1,81 @@ +import { describe, expect, it, vi, beforeEach } from 'vitest'; +import { buildExportOptions, buildExportPngOptions } from '../exportChatPresentation'; +import type { Message, ProviderConfig } from '@/types'; + +vi.mock('@/i18n', () => ({ + default: { + t: (key: string, opts?: { defaultValue?: string }) => { + const map: Record = { + 'chat.you': '你', + 'chat.assistant': '助手', + 'chat.system': '系统', + }; + return map[key] ?? opts?.defaultValue ?? key; + }, + }, +})); + +const providers = [ + { + id: 'p1', + name: 'OpenAI', + enabled: true, + models: [ + { + model_id: 'gpt-4o', + name: 'GPT-4o', + enabled: true, + }, + ], + }, +] as unknown as ProviderConfig[]; + +const theme = { + colorPrimary: '#1677ff', + colorPrimaryBg: '#e6f4ff', + colorPrimaryBorder: '#91caff', +}; + +describe('buildExportOptions', () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it('uses i18n role labels and profile name', () => { + const opts = buildExportOptions({ + userName: '小明', + theme, + providers, + }); + + expect(opts.roleLabels).toEqual({ + user: '你', + assistant: '助手', + system: '系统', + }); + expect(opts.userName).toBe('小明'); + }); + + it('resolves model labels for assistant titles without provider prefix', () => { + const opts = buildExportOptions({ + userName: '', + theme, + providers, + conversationModelId: 'gpt-4o', + conversationProviderId: 'p1', + }); + + const label = opts.getModelLabel?.({ + id: 'm1', + role: 'assistant', + model_id: 'gpt-4o', + provider_id: 'p1', + } as Message); + + expect(label).toBe('GPT-4o'); + }); + + it('keeps buildExportPngOptions as an alias', () => { + expect(buildExportPngOptions).toBe(buildExportOptions); + }); +}); diff --git a/src/lib/__tests__/filename.test.ts b/src/lib/__tests__/filename.test.ts new file mode 100644 index 00000000..197ab6cb --- /dev/null +++ b/src/lib/__tests__/filename.test.ts @@ -0,0 +1,44 @@ +import { describe, expect, it } from 'vitest'; +import { + formatExportFilenameTimestamp, + sanitizeExportFilename, + sanitizeFilenamePart, +} from '../filename'; + +describe('sanitizeFilenamePart', () => { + it('strips Windows-illegal characters', () => { + expect(sanitizeFilenamePart('关于:测试?文件*"name"|a')).toBe('关于-测试-文件-name-a'); + }); + + it('collapses whitespace and dashes', () => { + expect(sanitizeFilenamePart(' hello world -- chat ')).toBe('hello-world-chat'); + }); + + it('falls back when empty after sanitize', () => { + expect(sanitizeFilenamePart(':::')).toBe('aqbot'); + expect(sanitizeFilenamePart('***', 'chat')).toBe('chat'); + }); +}); + +describe('formatExportFilenameTimestamp', () => { + it('formats local time as YYYY-MM-DD_HHmmss', () => { + const date = new Date(2026, 2, 27, 14, 30, 52); // local March 27, 2026 14:30:52 + expect(formatExportFilenameTimestamp(date)).toBe('2026-03-27_143052'); + }); +}); + +describe('sanitizeExportFilename', () => { + it('appends extension, sanitizes title, and suffixes timestamp', () => { + const at = new Date(2026, 2, 27, 14, 30, 52); + expect(sanitizeExportFilename('对话: 计划?', 'png', 'chat', at)).toBe( + '对话-计划 - 2026-03-27_143052.png', + ); + }); + + it('normalizes extension without leading dot', () => { + const at = new Date(2026, 0, 1, 0, 0, 0); + expect(sanitizeExportFilename('notes', '.md', 'chat', at)).toBe( + 'notes - 2026-01-01_000000.md', + ); + }); +}); diff --git a/src/lib/chatImageActions.ts b/src/lib/chatImageActions.ts index 1a09d416..9daad9e9 100644 --- a/src/lib/chatImageActions.ts +++ b/src/lib/chatImageActions.ts @@ -1,5 +1,6 @@ import { invoke, isTauri } from './invoke'; import { Image as TauriImage } from '@tauri-apps/api/image'; +import { sanitizeFilenamePart as sanitizeFilenamePartShared } from './filename'; const IMAGE_MIME_EXTENSIONS: Record = { 'image/png': 'png', @@ -75,13 +76,7 @@ function getExtensionForImage(src: string, mimeType?: string | null) { } function sanitizeFilenamePart(value: string) { - const normalized = value - .trim() - .replace(/[\\/:*?"<>|]+/g, '-') - .replace(/\s+/g, '-') - .replace(/-+/g, '-') - .replace(/^-|-$/g, ''); - return normalized || 'aqbot-image'; + return sanitizeFilenamePartShared(value, 'aqbot-image'); } function ensureImageExtension(filename: string, src: string, mimeType?: string | null) { diff --git a/src/lib/exportChat.ts b/src/lib/exportChat.ts index e525522b..9337597b 100644 --- a/src/lib/exportChat.ts +++ b/src/lib/exportChat.ts @@ -1,5 +1,7 @@ import { isTauri } from '@/lib/invoke' -import { stripAqbotTags } from '@/lib/chatMarkdown' +import { stripAqbotTags, safeParseChatMarkdown, type ChatMarkdownNode } from '@/lib/chatMarkdown' +import { sanitizeExportFilename } from '@/lib/filename' +import { formatChatTime } from '@/components/chat/chatTime' import type { Message } from '@/types' function browserDownload(filename: string, content: string, mimeType: string) { @@ -59,8 +61,32 @@ async function writeToClipboard(text: string) { } } +export interface ExportRoleLabels { + user: string; + assistant: string; + system: string; +} + export interface TranscriptExportOptions { includeThinking?: boolean; + /** i18n role labels; defaults resolved from current language when omitted */ + roleLabels?: ExportRoleLabels; + /** Profile display name for user messages (falls back to roleLabels.user) */ + userName?: string; + /** Model display name for assistant messages (no provider prefix) */ + getModelLabel?: (message: Message) => string | undefined; + formatTime?: (createdAt: number) => string; +} + +export interface ExportPngTheme { + colorPrimary: string; + colorPrimaryBg: string; + colorPrimaryBorder: string; + colorFillSecondary?: string; +} + +export interface ExportMessagesPngOptions extends TranscriptExportOptions { + theme?: ExportPngTheme; } function getExportMessageContent(message: Message, options?: TranscriptExportOptions) { @@ -70,35 +96,398 @@ function getExportMessageContent(message: Message, options?: TranscriptExportOpt return message.content } +const FALLBACK_ROLE_LABELS: ExportRoleLabels = { + user: 'You', + assistant: 'Assistant', + system: 'System', +} + +function resolveRoleLabels(options?: TranscriptExportOptions): ExportRoleLabels { + return { + user: options?.roleLabels?.user || FALLBACK_ROLE_LABELS.user, + assistant: options?.roleLabels?.assistant || FALLBACK_ROLE_LABELS.assistant, + system: options?.roleLabels?.system || FALLBACK_ROLE_LABELS.system, + } +} + +/** Speaker line for export: user name / model name / system label (same as PNG). */ +export function resolveExportSpeakerLabel( + message: Message, + options?: TranscriptExportOptions, +): string { + const labels = resolveRoleLabels(options) + if (message.role === 'user') { + return options?.userName?.trim() || labels.user + } + if (message.role === 'system') { + return labels.system + } + return options?.getModelLabel?.(message) || labels.assistant +} + +function resolveExportMessageTime( + message: Message, + options?: TranscriptExportOptions, +): string | undefined { + if (message.created_at == null) return undefined + const formatTime = options?.formatTime ?? defaultFormatExportTime + return formatTime(message.created_at) +} + export function buildMarkdownTranscript(messages: Message[], title: string, options?: TranscriptExportOptions) { const lines: string[] = [`# ${title}`, ''] for (const m of messages) { - const role = m.role === 'user' ? 'User' : m.role === 'system' ? 'System' : 'Assistant' - lines.push(`## ${role}`, '', getExportMessageContent(m, options), '', '---', '') + const speaker = resolveExportSpeakerLabel(m, options) + lines.push(`## ${speaker}`, '', getExportMessageContent(m, options), '') + const time = resolveExportMessageTime(m, options) + if (time) lines.push(time, '') + lines.push('---', '') } return lines.join('\n') } export function buildTextTranscript(messages: Message[], title: string, options?: TranscriptExportOptions) { - const lines: string[] = [title, '='.repeat(title.length), ''] + const lines: string[] = [title, '='.repeat(Math.max(title.length, 1)), ''] for (const m of messages) { - const role = m.role === 'user' ? 'User' : m.role === 'system' ? 'System' : 'Assistant' - lines.push(`[${role}]`, '', getExportMessageContent(m, options), '', '---', '') + const speaker = resolveExportSpeakerLabel(m, options) + lines.push(`[${speaker}]`, '', getExportMessageContent(m, options), '') + const time = resolveExportMessageTime(m, options) + if (time) lines.push(time, '') + lines.push('---', '') } return lines.join('\n') } +// ── Export markdown → safe HTML ────────────────────────────────────── + +function escapeHtml(value: string): string { + return value + .replace(/&/g, '&') + .replace(//g, '>') + .replace(/"/g, '"') + .replace(/'/g, ''') +} + +function asRecord(node: ChatMarkdownNode): Record { + return node as unknown as Record +} + +function childrenOf(node: ChatMarkdownNode): ChatMarkdownNode[] { + const children = asRecord(node).children + return Array.isArray(children) ? children as ChatMarkdownNode[] : [] +} + +function renderInlineNodes(nodes: ChatMarkdownNode[]): string { + return nodes.map(renderInlineNode).join('') +} + +function renderInlineNode(node: ChatMarkdownNode): string { + const rec = asRecord(node) + switch (node.type) { + case 'text': + return escapeHtml(String(rec.content ?? rec.raw ?? '')) + case 'strong': + return `${renderInlineNodes(childrenOf(node))}` + case 'emphasis': + return `${renderInlineNodes(childrenOf(node))}` + case 'strikethrough': + return `${renderInlineNodes(childrenOf(node))}` + case 'highlight': + return `${renderInlineNodes(childrenOf(node))}` + case 'insert': + return `${renderInlineNodes(childrenOf(node))}` + case 'subscript': + return `${renderInlineNodes(childrenOf(node))}` + case 'superscript': + return `${renderInlineNodes(childrenOf(node))}` + case 'inline_code': + return `${escapeHtml(String(rec.code ?? ''))}` + case 'link': { + const href = String(rec.href ?? '') + const safeHref = /^(https?:|mailto:)/i.test(href) ? escapeHtml(href) : '#' + const text = childrenOf(node).length > 0 + ? renderInlineNodes(childrenOf(node)) + : escapeHtml(String(rec.text ?? href)) + return `${text}` + } + case 'image': { + const alt = escapeHtml(String(rec.alt ?? '')) + const src = String(rec.src ?? '') + // Only embed data URLs / relative media in export; remote src may break html2canvas. + if (/^(data:|https?:|http:\/\/aqbot-media\.localhost)/i.test(src)) { + return `${alt}` + } + return alt ? `[${alt}]` : '' + } + case 'hardbreak': + return '
' + case 'checkbox': + case 'checkbox_input': + return rec.checked ? '☑ ' : '☐ ' + case 'emoji': + return escapeHtml(String(rec.markup ?? rec.name ?? '')) + case 'math_inline': + return `${escapeHtml(String(rec.content ?? ''))}` + case 'footnote_reference': + return `[${escapeHtml(String(rec.id ?? ''))}]` + case 'html_inline': + // Never inject raw HTML; fall back to text children or escaped content. + if (childrenOf(node).length > 0) return renderInlineNodes(childrenOf(node)) + return escapeHtml(String(rec.content ?? '')) + case 'inline': + return renderInlineNodes(childrenOf(node)) + default: + if (childrenOf(node).length > 0) return renderInlineNodes(childrenOf(node)) + if (typeof rec.content === 'string') return escapeHtml(rec.content) + if (typeof rec.raw === 'string') return escapeHtml(rec.raw) + return '' + } +} + +function renderBlockNode(node: ChatMarkdownNode): string { + const rec = asRecord(node) + switch (node.type) { + case 'heading': { + const level = Math.min(6, Math.max(1, Number(rec.level) || 1)) + const inner = childrenOf(node).length > 0 + ? renderInlineNodes(childrenOf(node)) + : escapeHtml(String(rec.text ?? '')) + return `${inner}` + } + case 'paragraph': + return `

${renderInlineNodes(childrenOf(node))}

` + case 'blockquote': + return `
${childrenOf(node).map(renderBlockNode).join('')}
` + case 'list': { + const ordered = Boolean(rec.ordered) + const tag = ordered ? 'ol' : 'ul' + const start = ordered && typeof rec.start === 'number' && rec.start !== 1 + ? ` start="${rec.start}"` + : '' + const items = Array.isArray(rec.items) ? rec.items as ChatMarkdownNode[] : [] + const lis = items.map((item) => { + const itemChildren = childrenOf(item) + // List items may contain paragraphs or inlines. + const body = itemChildren.map((child) => { + if (child.type === 'paragraph' || child.type === 'list' || child.type === 'blockquote' || child.type === 'code_block') { + return renderBlockNode(child) + } + return renderInlineNode(child) + }).join('') + return `
  • ${body}
  • ` + }).join('') + return `<${tag} class="export-list"${start}>${lis}` + } + case 'code_block': { + const lang = escapeHtml(String(rec.language ?? '')) + const code = escapeHtml(String(rec.code ?? '')) + const langLabel = lang ? `
    ${lang}
    ` : '' + return `
    ${langLabel}
    ${code}
    ` + } + case 'thematic_break': + return '
    ' + case 'table': { + const header = rec.header as ChatMarkdownNode | undefined + const rows = Array.isArray(rec.rows) ? rec.rows as ChatMarkdownNode[] : [] + const renderRow = (row: ChatMarkdownNode, isHeader: boolean) => { + const cells = Array.isArray(asRecord(row).cells) + ? asRecord(row).cells as ChatMarkdownNode[] + : [] + const cellTag = isHeader ? 'th' : 'td' + const tds = cells.map((cell) => { + const align = asRecord(cell).align + const alignAttr = align === 'left' || align === 'right' || align === 'center' + ? ` style="text-align:${align}"` + : '' + return `<${cellTag}${alignAttr}>${renderInlineNodes(childrenOf(cell))}` + }).join('') + return `${tds}` + } + const thead = header ? `${renderRow(header, true)}` : '' + const tbody = `${rows.map((r) => renderRow(r, false)).join('')}` + return `${thead}${tbody}
    ` + } + case 'math_block': + return `
    ${escapeHtml(String(rec.content ?? ''))}
    ` + case 'admonition': { + const title = escapeHtml(String(rec.title || rec.kind || 'Note')) + return `
    ${title}
    ${childrenOf(node).map(renderBlockNode).join('')}
    ` + } + case 'html_block': + case 'html-render': + // Strip/escape untrusted HTML blocks. + return `
    ${escapeHtml(String(rec.content ?? rec.raw ?? ''))}
    ` + case 'think': + case 'web-search': + case 'web-search-query': + case 'knowledge-retrieval': + case 'memory-retrieval': + case 'tool-call': + return '' // stripped for clean share images + default: { + // Custom components / unknown: try children, else plain escaped raw. + if (childrenOf(node).length > 0) { + return childrenOf(node).map((child) => { + if ( + child.type === 'paragraph' + || child.type === 'heading' + || child.type === 'list' + || child.type === 'code_block' + || child.type === 'blockquote' + || child.type === 'table' + ) { + return renderBlockNode(child) + } + return renderInlineNode(child) + }).join('') + } + if (typeof rec.content === 'string' && rec.content.trim()) { + return `

    ${escapeHtml(rec.content)}

    ` + } + if (typeof rec.raw === 'string' && rec.raw.trim()) { + return `

    ${escapeHtml(rec.raw)}

    ` + } + return '' + } + } +} + +/** Convert chat markdown to safe HTML for PNG export (XSS-escaped, no raw HTML passthrough). */ +export function renderExportMarkdownHtml(content: string): string { + const text = content.trim() + if (!text) return '' + try { + const nodes = safeParseChatMarkdown(text) + return nodes.map(renderBlockNode).join('') + } catch (error) { + console.error('Export markdown render failed, falling back to plain text:', error) + return `

    ${escapeHtml(text)}

    ` + } +} + +const EXPORT_MARKDOWN_CSS = ` +.export-msg-name { + font-size: 12px; + font-weight: 600; + color: #6b7280; + line-height: 1.4; + margin: 0 0 8px; + padding: 0; +} +.export-md { font-size: 14px; line-height: 1.65; word-break: break-word; color: #111827; } +.export-md > :first-child { margin-top: 0; } +.export-md > :last-child { margin-bottom: 0; } +.export-h { margin: 12px 0 8px; font-weight: 600; line-height: 1.35; color: #111827; } +.export-h:first-child { margin-top: 0; } +h1.export-h { font-size: 1.35em; } +h2.export-h { font-size: 1.2em; } +h3.export-h { font-size: 1.08em; } +h4.export-h, h5.export-h, h6.export-h { font-size: 1em; } +.export-p { margin: 0 0 10px; white-space: pre-wrap; } +.export-p:last-child { margin-bottom: 0; } +.export-list { margin: 0 0 10px; padding-left: 1.4em; } +.export-list li { margin: 2px 0; } +.export-list li > .export-p { margin: 0; white-space: normal; } +.export-quote { + margin: 0 0 10px; + padding: 6px 12px; + border-left: 3px solid #93c5fd; + background: #f8fafc; + color: #374151; +} +.export-code-inline { + font-family: ui-monospace, SFMono-Regular, Menlo, Monaco, Consolas, monospace; + font-size: 0.9em; + background: #f3f4f6; + border: 1px solid #e5e7eb; + border-radius: 4px; + padding: 0 4px; +} +.export-code-block { + margin: 0 0 10px; + border: 1px solid #e5e7eb; + border-radius: 8px; + overflow: hidden; + background: #f9fafb; +} +.export-code-lang { + font-size: 11px; + color: #6b7280; + padding: 4px 10px; + border-bottom: 1px solid #e5e7eb; + background: #f3f4f6; +} +.export-code-block pre { + margin: 0; + padding: 10px 12px; + overflow-x: auto; + font-family: ui-monospace, SFMono-Regular, Menlo, Monaco, Consolas, monospace; + font-size: 12.5px; + line-height: 1.5; + white-space: pre; +} +.export-code-block code { font-family: inherit; background: none; border: none; padding: 0; } +.export-table { + width: 100%; + border-collapse: collapse; + margin: 0 0 10px; + font-size: 13px; +} +.export-table th, .export-table td { + border: 1px solid #e5e7eb; + padding: 6px 8px; + text-align: left; + vertical-align: top; +} +.export-table th { background: #f3f4f6; font-weight: 600; } +.export-hr { border: none; border-top: 1px solid #e5e7eb; margin: 12px 0; } +.export-math, .export-math-block { + font-family: ui-monospace, SFMono-Regular, Menlo, Monaco, Consolas, monospace; + font-size: 0.92em; + background: #f8fafc; +} +.export-math-block { margin: 0 0 10px; padding: 8px 10px; border-radius: 6px; white-space: pre-wrap; } +.export-admonition { + margin: 0 0 10px; + padding: 8px 10px; + border: 1px solid #e5e7eb; + border-radius: 8px; + background: #f9fafb; +} +.export-admonition-title { font-weight: 600; margin-bottom: 4px; font-size: 12px; color: #4b5563; } +.export-img { max-width: 100%; height: auto; border-radius: 6px; display: block; margin: 6px 0; } +.export-img-fallback { color: #6b7280; font-size: 12px; } +.export-html-fallback { + margin: 0 0 10px; + padding: 8px; + background: #f9fafb; + border-radius: 6px; + font-size: 12px; + white-space: pre-wrap; + color: #4b5563; +} +.export-md a { color: #2563eb; text-decoration: underline; } +` + +function injectExportMarkdownStyles(host: HTMLElement) { + const style = document.createElement('style') + style.textContent = EXPORT_MARKDOWN_CSS + host.appendChild(style) +} + async function canvasToPngFile(canvas: HTMLCanvasElement, title: string) { + const safeName = sanitizeExportFilename(title, 'png', 'chat') if (isTauri()) { const blob = await new Promise((resolve) => canvas.toBlob(resolve, 'image/png')) if (!blob) return false const buffer = new Uint8Array(await blob.arrayBuffer()) - return saveFile(`${title}.png`, buffer, [{ name: 'PNG Image', extensions: ['png'] }]) + return saveFile(safeName, buffer, [{ name: 'PNG Image', extensions: ['png'] }]) } // Browser fallback const link = document.createElement('a') - link.download = `${title}.png` + link.download = safeName link.href = canvas.toDataURL('image/png') link.click() return true @@ -153,17 +542,44 @@ export async function exportAsPNG(element: HTMLElement | null, title: string) { return canvasToPngFile(canvas, title) } +async function resolveDefaultRoleLabels(): Promise { + try { + const i18n = (await import('@/i18n')).default + return { + user: i18n.t('chat.you', { defaultValue: 'You' }), + assistant: i18n.t('chat.assistant', { defaultValue: 'Assistant' }), + system: i18n.t('chat.system', { defaultValue: 'System' }), + } + } catch { + return { ...FALLBACK_ROLE_LABELS } + } +} + +function defaultFormatExportTime(createdAt: number): string { + return formatChatTime(createdAt) +} + /** * Render selected messages into a clean off-screen card layout, then capture PNG. - * Avoids viewport clipping and action-icon layout bugs from the live chat DOM. + * No avatars — role line is user name / model name / system label; footer is time only. */ export async function exportMessagesAsPNG( messages: Message[], title: string, - options?: TranscriptExportOptions, + options?: ExportMessagesPngOptions, ) { if (messages.length === 0) return false + const roleLabels = options?.roleLabels ?? await resolveDefaultRoleLabels() + const theme: ExportPngTheme = { + colorPrimary: options?.theme?.colorPrimary ?? '#1677ff', + colorPrimaryBg: options?.theme?.colorPrimaryBg ?? '#e6f4ff', + colorPrimaryBorder: options?.theme?.colorPrimaryBorder ?? '#91caff', + colorFillSecondary: options?.theme?.colorFillSecondary ?? '#f3f4f6', + } + const userName = options?.userName?.trim() || roleLabels.user || 'You' + const formatTime = options?.formatTime ?? defaultFormatExportTime + const host = document.createElement('div') host.setAttribute('data-export-share-root', 'true') host.style.cssText = [ @@ -178,6 +594,8 @@ export async function exportMessagesAsPNG( 'box-sizing:border-box', ].join(';') + injectExportMarkdownStyles(host) + const heading = document.createElement('div') heading.style.cssText = 'font-size:18px;font-weight:600;margin:0 0 4px;line-height:1.4;' heading.textContent = title @@ -189,26 +607,43 @@ export async function exportMessagesAsPNG( host.appendChild(meta) for (const message of messages) { - const card = document.createElement('div') const isUser = message.role === 'user' + const isSystem = message.role === 'system' + const card = document.createElement('div') card.style.cssText = [ 'margin:0 0 14px', 'padding:12px 14px', 'border-radius:12px', - `background:${isUser ? '#eff6ff' : '#f9fafb'}`, - `border:1px solid ${isUser ? '#dbeafe' : '#e5e7eb'}`, + `background:${isUser ? theme.colorPrimaryBg : '#f9fafb'}`, + `border:1px solid ${isUser ? theme.colorPrimaryBorder : '#e5e7eb'}`, ].join(';') - const role = document.createElement('div') - role.style.cssText = 'font-size:12px;font-weight:600;color:#6b7280;margin:0 0 8px;' - role.textContent = isUser ? 'User' : message.role === 'system' ? 'System' : 'Assistant' - card.appendChild(role) + // Title: user name / model name / system label (no avatar) + const nameEl = document.createElement('div') + nameEl.className = 'export-msg-name' + nameEl.style.cssText = 'font-size:12px;font-weight:600;color:#6b7280;line-height:1.4;margin:0 0 8px;' + if (isUser) { + nameEl.textContent = userName + } else if (isSystem) { + nameEl.textContent = roleLabels.system + } else { + nameEl.textContent = options?.getModelLabel?.(message) || roleLabels.assistant + } + card.appendChild(nameEl) const body = document.createElement('div') - body.style.cssText = 'font-size:14px;line-height:1.65;white-space:pre-wrap;word-break:break-word;' - body.textContent = getExportMessageContent(message, options) + body.className = 'export-md' + body.innerHTML = renderExportMarkdownHtml(getExportMessageContent(message, options)) card.appendChild(body) + // Footer: time only + if (message.created_at != null) { + const footer = document.createElement('div') + footer.style.cssText = 'margin-top:8px;font-size:11px;line-height:1.4;color:#9ca3af;' + footer.textContent = formatTime(message.created_at) + card.appendChild(footer) + } + host.appendChild(card) } @@ -234,12 +669,18 @@ export function buildJsonTranscript(messages: Message[], title: string, options? const data = { title, exported_at: new Date().toISOString(), - messages: messages.map((m) => ({ - role: m.role, - content: getExportMessageContent(m, options), - ...(options?.includeThinking === false ? {} : { thinking: m.thinking }), - created_at: m.created_at, - })), + messages: messages.map((m) => { + const time = resolveExportMessageTime(m, options) + return { + role: m.role, + // Display name: user profile / model name / system label (same as PNG/md/txt) + name: resolveExportSpeakerLabel(m, options), + content: getExportMessageContent(m, options), + ...(options?.includeThinking === false ? {} : { thinking: m.thinking }), + ...(time ? { time } : {}), + created_at: m.created_at, + } + }), } return JSON.stringify(data, null, 2) } @@ -258,13 +699,13 @@ export async function copyTranscript( } export async function exportAsMarkdown(messages: Message[], title: string, options?: TranscriptExportOptions) { - return saveFile(`${title}.md`, buildMarkdownTranscript(messages, title, options), [{ name: 'Markdown', extensions: ['md'] }]) + return saveFile(sanitizeExportFilename(title, 'md', 'chat'), buildMarkdownTranscript(messages, title, options), [{ name: 'Markdown', extensions: ['md'] }]) } export async function exportAsText(messages: Message[], title: string, options?: TranscriptExportOptions) { - return saveFile(`${title}.txt`, buildTextTranscript(messages, title, options), [{ name: 'Text', extensions: ['txt'] }]) + return saveFile(sanitizeExportFilename(title, 'txt', 'chat'), buildTextTranscript(messages, title, options), [{ name: 'Text', extensions: ['txt'] }]) } export async function exportAsJSON(messages: Message[], title: string, options?: TranscriptExportOptions) { - return saveFile(`${title}.json`, buildJsonTranscript(messages, title, options), [{ name: 'JSON', extensions: ['json'] }]) + return saveFile(sanitizeExportFilename(title, 'json', 'chat'), buildJsonTranscript(messages, title, options), [{ name: 'JSON', extensions: ['json'] }]) } diff --git a/src/lib/exportChatPresentation.ts b/src/lib/exportChatPresentation.ts new file mode 100644 index 00000000..a33d61ac --- /dev/null +++ b/src/lib/exportChatPresentation.ts @@ -0,0 +1,79 @@ +import i18n from '@/i18n' +import { formatChatTime } from '@/components/chat/chatTime' +import type { Message, ProviderConfig } from '@/types' +import type { ExportMessagesPngOptions, ExportPngTheme, TranscriptExportOptions } from './exportChat' + +export type ExportPresentationInput = { + /** Display name for user messages; empty falls back to i18n "you" */ + userName?: string | null + theme: ExportPngTheme + providers: ProviderConfig[] + /** Conversation-level fallback when message.model_id is empty */ + conversationModelId?: string | null + conversationProviderId?: string | null +} + +function resolveModelLabel( + message: Message, + providers: ProviderConfig[], + conversationModelId?: string | null, + conversationProviderId?: string | null, +): string | undefined { + const mid = message.model_id ?? conversationModelId + const pid = message.provider_id ?? conversationProviderId + if (!mid) return undefined + const provider = pid ? providers.find((p) => p.id === pid) : undefined + const model = provider?.models.find((m) => m.model_id === mid) + // Title uses model name only (no provider/platform prefix). + if (model?.name) return model.name + for (const p of providers) { + const m = p.models.find((item) => item.model_id === mid) + if (m?.name) return m.name + } + return mid +} + +/** Shared presentation options for PNG / Markdown / Text / copy (i18n, model name, time). */ +export function buildExportOptions(input: ExportPresentationInput): ExportMessagesPngOptions { + const { userName, theme, providers, conversationModelId, conversationProviderId } = input + + return { + roleLabels: { + user: i18n.t('chat.you', { defaultValue: 'You' }), + assistant: i18n.t('chat.assistant', { defaultValue: 'Assistant' }), + system: i18n.t('chat.system', { defaultValue: 'System' }), + }, + userName: userName?.trim() || undefined, + theme, + getModelLabel: (message) => resolveModelLabel( + message, + providers, + conversationModelId, + conversationProviderId, + ), + formatTime: (createdAt) => formatChatTime(createdAt), + } +} + +/** @deprecated Prefer buildExportOptions — same implementation. */ +export const buildExportPngOptions = buildExportOptions + +/** Transcript-only subset (no theme) for callers that only need md/text. */ +export function buildTranscriptExportOptions( + input: Omit & { theme?: ExportPngTheme }, +): TranscriptExportOptions { + const full = buildExportOptions({ + ...input, + theme: input.theme ?? { + colorPrimary: '#1677ff', + colorPrimaryBg: '#e6f4ff', + colorPrimaryBorder: '#91caff', + }, + }) + return { + roleLabels: full.roleLabels, + userName: full.userName, + getModelLabel: full.getModelLabel, + formatTime: full.formatTime, + } +} diff --git a/src/lib/filename.ts b/src/lib/filename.ts new file mode 100644 index 00000000..a86a652f --- /dev/null +++ b/src/lib/filename.ts @@ -0,0 +1,45 @@ +/** + * Sanitize a user-facing string for use as a file name segment. + * Strips Windows/macOS/Linux illegal path characters and collapses whitespace. + */ +export function sanitizeFilenamePart(value: string, fallback = 'aqbot'): string { + const normalized = value + .trim() + .replace(/[\\/:*?"<>|]+/g, '-') + .replace(/\s+/g, '-') + .replace(/-+/g, '-') + .replace(/^-|-$/g, ''); + return normalized || fallback; +} + +/** Filesystem-safe local timestamp: `2026-03-27_143052` */ +export function formatExportFilenameTimestamp(date: Date = new Date()): string { + const pad = (n: number) => String(n).padStart(2, '0'); + return [ + date.getFullYear(), + '-', + pad(date.getMonth() + 1), + '-', + pad(date.getDate()), + '_', + pad(date.getHours()), + pad(date.getMinutes()), + pad(date.getSeconds()), + ].join(''); +} + +/** + * Build a safe download file name with extension and a unique timestamp suffix. + * Example: `对话-计划 - 2026-03-27_143052.png` + */ +export function sanitizeExportFilename( + title: string, + extension: string, + fallback = 'chat', + at: Date = new Date(), +): string { + const base = sanitizeFilenamePart(title, fallback); + const ext = extension.replace(/^\./, '').toLowerCase() || 'bin'; + const timestamp = formatExportFilenameTimestamp(at); + return `${base} - ${timestamp}.${ext}`; +} From 74bc08587841bb4a95e1e70bd1a119feb6f3315e Mon Sep 17 00:00:00 2001 From: licoy Date: Fri, 7 Aug 2026 17:32:02 +0800 Subject: [PATCH 003/108] chore(version): bump version to v0.0.118 --- package.json | 2 +- src-tauri/tauri.conf.json | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/package.json b/package.json index ea8a0ed1..4a8a6402 100644 --- a/package.json +++ b/package.json @@ -1,7 +1,7 @@ { "name": "aqbot", "private": true, - "version": "0.0.117", + "version": "0.0.118", "license": "AGPL-3.0-only", "packageManager": "pnpm@10.32.1", "type": "module", diff --git a/src-tauri/tauri.conf.json b/src-tauri/tauri.conf.json index 6595e832..7934d78a 100644 --- a/src-tauri/tauri.conf.json +++ b/src-tauri/tauri.conf.json @@ -1,7 +1,7 @@ { "$schema": "https://schema.tauri.app/config/2", "productName": "AQBot", - "version": "0.0.117", + "version": "0.0.118", "identifier": "top.aqbot.desktop", "build": { "beforeDevCommand": "pnpm dev", From 4df00d22b572702dd38c3273acef0df81f6b3392 Mon Sep 17 00:00:00 2001 From: licoy Date: Fri, 7 Aug 2026 18:06:10 +0800 Subject: [PATCH 004/108] =?UTF-8?q?feat(chat):=20=E5=A2=9E=E5=BC=BA?= =?UTF-8?q?=E4=B8=8A=E4=B8=8B=E6=96=87=E5=8E=8B=E7=BC=A9=E4=BF=9D=E7=95=99?= =?UTF-8?q?=E5=B0=BE=E9=83=A8=E6=B6=88=E6=81=AF=E4=B8=8E=E5=8E=9F=E6=96=87?= =?UTF-8?q?=E9=87=8D=E8=AF=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 说明: - 压缩时支持保留最近 N 条消息不进摘要,全局默认可配置,会话可覆盖 - 持久化压缩输入原文,摘要弹窗支持查看原文并用当前全局提示词重试 - 压缩提示词在提交正文末尾复述约束,降低模型顺着原文续写的概率 - 上下文 token 圈增加立即压缩与查看摘要入口 操作: - 设置 → 对话设置:配置全局「压缩时保留最近消息数」 - 对话设置弹窗:可覆盖 keep-last-N - 压缩模型与提示词仍在设置 → 默认模型中配置(留空用系统默认) - 升级后自动 migration:conversations.compression_keep_last_n、conversation_summaries.source_text Closes #119 --- .../core/src/entity/conversation_summaries.rs | 3 + .../crates/core/src/entity/conversations.rs | 3 + .../crates/core/src/repo/chatgpt_import.rs | 1 + .../crates/core/src/repo/cherry_import.rs | 1 + .../crates/core/src/repo/conversation.rs | 40 ++ .../crates/core/src/repo/kelivo_import.rs | 1 + src-tauri/crates/core/src/repo/message.rs | 4 +- src-tauri/crates/core/src/types.rs | 14 + .../crates/core/tests/repo_integration.rs | 1 + src-tauri/crates/migration/src/lib.rs | 2 + ...0808_000001_compression_keep_and_source.rs | 59 +++ src-tauri/src/commands/agent.rs | 1 + src-tauri/src/commands/conversations.rs | 456 ++++++++++++++---- src-tauri/src/context_manager.rs | 327 ++++++++----- src-tauri/src/lib.rs | 1 + src/components/chat/ChatView.tsx | 142 +++++- .../chat/ConversationSettingsModal.tsx | 45 +- src/components/chat/InputArea.tsx | 56 ++- .../chat/__tests__/ModelSelector.test.tsx | 1 + .../__tests__/conversationListModel.test.ts | 1 + .../settings/ConversationSettings.tsx | 42 +- .../__tests__/ConversationSettings.test.tsx | 20 + .../__tests__/DefaultModelSettings.test.tsx | 1 + src/i18n/locales/en-US.json | 15 + src/i18n/locales/zh-CN.json | 15 + src/lib/browserMock.ts | 1 + src/stores/conversationStore.ts | 33 ++ src/stores/settingsStore.ts | 1 + src/types/index.ts | 11 + tests/performance/browserFixture.ts | 1 + 30 files changed, 1059 insertions(+), 240 deletions(-) create mode 100644 src-tauri/crates/migration/src/m20260808_000001_compression_keep_and_source.rs diff --git a/src-tauri/crates/core/src/entity/conversation_summaries.rs b/src-tauri/crates/core/src/entity/conversation_summaries.rs index 8751888f..e76a6458 100644 --- a/src-tauri/crates/core/src/entity/conversation_summaries.rs +++ b/src-tauri/crates/core/src/entity/conversation_summaries.rs @@ -9,6 +9,9 @@ pub struct Model { pub conversation_id: String, pub summary_text: String, pub compressed_until_message_id: Option, + /// Rendered compression input (conversation body + optional prior summary) + /// used for display and retry. `None` for summaries created before this field. + pub source_text: Option, pub token_count: Option, pub model_used: Option, pub created_at: i64, diff --git a/src-tauri/crates/core/src/entity/conversations.rs b/src-tauri/crates/core/src/entity/conversations.rs index bb151235..38f18511 100644 --- a/src-tauri/crates/core/src/entity/conversations.rs +++ b/src-tauri/crates/core/src/entity/conversations.rs @@ -34,6 +34,9 @@ pub struct Model { /// Max provider history messages for this conversation. `None` falls back /// to the global `default_context_count` setting. Values ≥ 50 mean unlimited. pub context_message_limit: Option, + /// When compressing context, keep the last N compressible messages in clear + /// text (not included in the summary). `None` means default (3). `0` keeps none. + pub compression_keep_last_n: Option, pub category_id: Option, pub parent_conversation_id: Option, pub mode: String, diff --git a/src-tauri/crates/core/src/repo/chatgpt_import.rs b/src-tauri/crates/core/src/repo/chatgpt_import.rs index 9aa2c0c7..cd7d0bf2 100644 --- a/src-tauri/crates/core/src/repo/chatgpt_import.rs +++ b/src-tauri/crates/core/src/repo/chatgpt_import.rs @@ -201,6 +201,7 @@ pub async fn import_chatgpt_export_from_path( research_mode: Set(0), context_compression: Set(0), context_message_limit: Set(None), + compression_keep_last_n: Set(None), category_id: Set(None), parent_conversation_id: Set(None), mode: Set("chat".to_string()), diff --git a/src-tauri/crates/core/src/repo/cherry_import.rs b/src-tauri/crates/core/src/repo/cherry_import.rs index 8174db7c..4a2c360d 100644 --- a/src-tauri/crates/core/src/repo/cherry_import.rs +++ b/src-tauri/crates/core/src/repo/cherry_import.rs @@ -456,6 +456,7 @@ pub async fn import_cherry_studio_backup_from_path_with_root( research_mode: Set(0), context_compression: Set(0), context_message_limit: Set(None), + compression_keep_last_n: Set(None), category_id: Set(None), parent_conversation_id: Set(None), mode: Set("chat".to_string()), diff --git a/src-tauri/crates/core/src/repo/conversation.rs b/src-tauri/crates/core/src/repo/conversation.rs index b58b0a27..9fa1b161 100644 --- a/src-tauri/crates/core/src/repo/conversation.rs +++ b/src-tauri/crates/core/src/repo/conversation.rs @@ -34,6 +34,7 @@ fn conversation_from_entity(m: conversations::Model) -> Conversation { is_archived: m.is_archived != 0, context_compression: m.context_compression != 0, context_message_limit: m.context_message_limit.map(|v| v as u32), + compression_keep_last_n: m.compression_keep_last_n.map(|v| v as u32), category_id: m.category_id, parent_conversation_id: m.parent_conversation_id, mode: m.mode, @@ -211,6 +212,9 @@ pub async fn update_conversation( if let Some(context_message_limit) = input.context_message_limit { am.context_message_limit = Set(context_message_limit.map(|v| v as i32)); } + if let Some(compression_keep_last_n) = input.compression_keep_last_n { + am.compression_keep_last_n = Set(compression_keep_last_n.map(|v| v as i32)); + } if let Some(category_id) = input.category_id { am.category_id = Set(category_id); } @@ -509,6 +513,7 @@ pub async fn branch_conversation( is_archived: Set(0), context_compression: Set(source.context_compression), context_message_limit: Set(source.context_message_limit), + compression_keep_last_n: Set(source.compression_keep_last_n), category_id: Set(source.category_id.clone()), parent_conversation_id: Set(parent_id), research_mode: Set(source.research_mode), @@ -722,6 +727,7 @@ fn summary_from_entity(m: conversation_summaries::Model) -> ConversationSummary conversation_id: m.conversation_id, summary_text: m.summary_text, compressed_until_message_id: m.compressed_until_message_id, + source_text: m.source_text, token_count: m.token_count.map(|v| v as u32), model_used: m.model_used, created_at: m.created_at, @@ -749,6 +755,7 @@ pub async fn upsert_summary( compressed_until_message_id: Option<&str>, token_count: Option, model_used: Option<&str>, + source_text: Option<&str>, ) -> Result { let now = now_ts(); @@ -765,6 +772,9 @@ pub async fn upsert_summary( Set(compressed_until_message_id.map(|s| s.to_string())); am.token_count = Set(token_count.map(|v| v as i64)); am.model_used = Set(model_used.map(|s| s.to_string())); + if let Some(source) = source_text { + am.source_text = Set(Some(source.to_string())); + } am.updated_at = Set(now); am.update(db).await?; } @@ -777,6 +787,7 @@ pub async fn upsert_summary( compressed_until_message_id: Set( compressed_until_message_id.map(|s| s.to_string()), ), + source_text: Set(source_text.map(|s| s.to_string())), token_count: Set(token_count.map(|v| v as i64)), model_used: Set(model_used.map(|s| s.to_string())), created_at: Set(now), @@ -794,6 +805,35 @@ pub async fn upsert_summary( }) } +/// Update only the summary text/metadata after a retry, preserving boundary and source. +pub async fn update_summary_text( + db: &DatabaseConnection, + conversation_id: &str, + summary_text: &str, + token_count: Option, + model_used: Option<&str>, +) -> Result { + let now = now_ts(); + let existing = conversation_summaries::Entity::find() + .filter(conversation_summaries::Column::ConversationId.eq(conversation_id)) + .one(db) + .await? + .ok_or_else(|| AQBotError::NotFound(format!("Summary for conversation {}", conversation_id)))?; + + let mut am: conversation_summaries::ActiveModel = existing.into(); + am.summary_text = Set(summary_text.to_string()); + am.token_count = Set(token_count.map(|v| v as i64)); + am.model_used = Set(model_used.map(|s| s.to_string())); + am.updated_at = Set(now); + am.update(db).await?; + + get_summary(db, conversation_id).await?.ok_or_else(|| { + AQBotError::Database(sea_orm::DbErr::Custom( + "Failed to read back updated summary".into(), + )) + }) +} + pub async fn delete_summary(db: &DatabaseConnection, conversation_id: &str) -> Result<()> { conversation_summaries::Entity::delete_many() .filter(conversation_summaries::Column::ConversationId.eq(conversation_id)) diff --git a/src-tauri/crates/core/src/repo/kelivo_import.rs b/src-tauri/crates/core/src/repo/kelivo_import.rs index 27df57f2..1c7dfe0d 100644 --- a/src-tauri/crates/core/src/repo/kelivo_import.rs +++ b/src-tauri/crates/core/src/repo/kelivo_import.rs @@ -394,6 +394,7 @@ pub async fn import_kelivo_backup_from_path_with_root( research_mode: Set(0), context_compression: Set(0), context_message_limit: Set(None), + compression_keep_last_n: Set(None), category_id: Set(None), parent_conversation_id: Set(None), mode: Set("chat".to_string()), diff --git a/src-tauri/crates/core/src/repo/message.rs b/src-tauri/crates/core/src/repo/message.rs index d03698d8..dfaa8e27 100644 --- a/src-tauri/crates/core/src/repo/message.rs +++ b/src-tauri/crates/core/src/repo/message.rs @@ -1449,7 +1449,7 @@ mod tests { ) .await .unwrap(); - conversation::upsert_summary(db, &conv.id, "old summary", Some(&active.id), Some(12), Some("model-1")) + conversation::upsert_summary(db, &conv.id, "old summary", Some(&active.id), Some(12), Some("model-1"), None) .await .unwrap(); let later_user = create_message(db, &conv.id, MessageRole::User, "later", &[], None, 0) @@ -1535,7 +1535,7 @@ mod tests { .await .unwrap(); set_conversation_active_message_count(db, &conv.id, 2).await.unwrap(); - conversation::upsert_summary(db, &conv.id, "summary", Some(&assistant.id), Some(5), Some("model-1")) + conversation::upsert_summary(db, &conv.id, "summary", Some(&assistant.id), Some(5), Some("model-1"), None) .await .unwrap(); diff --git a/src-tauri/crates/core/src/types.rs b/src-tauri/crates/core/src/types.rs index b1c2735f..29358050 100644 --- a/src-tauri/crates/core/src/types.rs +++ b/src-tauri/crates/core/src/types.rs @@ -585,6 +585,9 @@ pub struct Conversation { /// Per-conversation cap on history messages sent to the model. /// `None` falls back to global `default_context_count`. Values ≥ 50 mean unlimited. pub context_message_limit: Option, + /// Keep the last N compressible messages out of compression. + /// `None` uses the default (3). `Some(0)` keeps none (compress all eligible). + pub compression_keep_last_n: Option, pub category_id: Option, pub parent_conversation_id: Option, pub mode: String, @@ -700,6 +703,9 @@ pub struct ConversationSummary { pub conversation_id: String, pub summary_text: String, pub compressed_until_message_id: Option, + /// Compression input text (for viewing and retry). Absent on legacy rows. + #[serde(default)] + pub source_text: Option, pub token_count: Option, pub model_used: Option, pub created_at: i64, @@ -736,6 +742,9 @@ pub struct UpdateConversationInput { /// Set to `Some(None)` to clear the override (use global default). #[serde(default, deserialize_with = "deserialize_double_option")] pub context_message_limit: Option>, + /// Set to `Some(None)` to clear and use the default keep-last-N (3). + #[serde(default, deserialize_with = "deserialize_double_option")] + pub compression_keep_last_n: Option>, #[serde(default, deserialize_with = "deserialize_double_option")] pub category_id: Option>, #[serde(default, deserialize_with = "deserialize_double_option")] @@ -1601,6 +1610,10 @@ pub struct AppSettings { pub compression_top_p: Option, pub compression_frequency_penalty: Option, pub compression_prompt: Option, + /// Global default for how many trailing messages to keep clear when compressing. + /// Per-conversation `compression_keep_last_n` overrides this. `None` → 3. + #[serde(default)] + pub default_compression_keep_last_n: Option, /// Model metadata source. Built-in is offline and is the default. pub model_catalog_source: ModelCatalogSourcePreference, pub proxy_type: Option, @@ -1762,6 +1775,7 @@ impl Default for AppSettings { compression_top_p: None, compression_frequency_penalty: None, compression_prompt: None, + default_compression_keep_last_n: None, model_catalog_source: ModelCatalogSourcePreference::Builtin, proxy_type: None, proxy_address: None, diff --git a/src-tauri/crates/core/tests/repo_integration.rs b/src-tauri/crates/core/tests/repo_integration.rs index df00472c..c83fb25b 100644 --- a/src-tauri/crates/core/tests/repo_integration.rs +++ b/src-tauri/crates/core/tests/repo_integration.rs @@ -349,6 +349,7 @@ async fn test_conversation_update_input() { enabled_memory_namespace_ids: Some(vec!["mem-a".into()]), context_compression: None, context_message_limit: Some(Some(3)), + compression_keep_last_n: None, category_id: None, parent_conversation_id: None, mode: None, diff --git a/src-tauri/crates/migration/src/lib.rs b/src-tauri/crates/migration/src/lib.rs index a2f09a76..7eb31b63 100644 --- a/src-tauri/crates/migration/src/lib.rs +++ b/src-tauri/crates/migration/src/lib.rs @@ -42,6 +42,7 @@ mod m20260724_000001_add_model_metadata; mod m20260725_000001_add_provider_aws_region; mod m20260806_000001_add_role_capability_bindings; mod m20260807_000001_add_conversation_context_message_limit; +mod m20260808_000001_compression_keep_and_source; pub struct Migrator; @@ -91,6 +92,7 @@ impl MigratorTrait for Migrator { Box::new(m20260725_000001_add_provider_aws_region::Migration), Box::new(m20260806_000001_add_role_capability_bindings::Migration), Box::new(m20260807_000001_add_conversation_context_message_limit::Migration), + Box::new(m20260808_000001_compression_keep_and_source::Migration), ] } } diff --git a/src-tauri/crates/migration/src/m20260808_000001_compression_keep_and_source.rs b/src-tauri/crates/migration/src/m20260808_000001_compression_keep_and_source.rs new file mode 100644 index 00000000..256bc2ba --- /dev/null +++ b/src-tauri/crates/migration/src/m20260808_000001_compression_keep_and_source.rs @@ -0,0 +1,59 @@ +use sea_orm_migration::prelude::*; + +#[derive(DeriveMigrationName)] +pub struct Migration; + +#[async_trait::async_trait] +impl MigrationTrait for Migration { + async fn up(&self, manager: &SchemaManager) -> Result<(), DbErr> { + manager + .alter_table( + Table::alter() + .table(Alias::new("conversations")) + .add_column( + ColumnDef::new(Alias::new("compression_keep_last_n")) + .integer() + .null(), + ) + .to_owned(), + ) + .await?; + + manager + .alter_table( + Table::alter() + .table(Alias::new("conversation_summaries")) + .add_column( + ColumnDef::new(Alias::new("source_text")) + .text() + .null(), + ) + .to_owned(), + ) + .await?; + + Ok(()) + } + + async fn down(&self, manager: &SchemaManager) -> Result<(), DbErr> { + manager + .alter_table( + Table::alter() + .table(Alias::new("conversation_summaries")) + .drop_column(Alias::new("source_text")) + .to_owned(), + ) + .await?; + + manager + .alter_table( + Table::alter() + .table(Alias::new("conversations")) + .drop_column(Alias::new("compression_keep_last_n")) + .to_owned(), + ) + .await?; + + Ok(()) + } +} diff --git a/src-tauri/src/commands/agent.rs b/src-tauri/src/commands/agent.rs index 8e7240a5..6a4a62e3 100644 --- a/src-tauri/src/commands/agent.rs +++ b/src-tauri/src/commands/agent.rs @@ -2295,6 +2295,7 @@ mod tests { is_archived: false, context_compression: false, context_message_limit: None, + compression_keep_last_n: None, category_id: None, parent_conversation_id: None, mode: "agent".to_string(), diff --git a/src-tauri/src/commands/conversations.rs b/src-tauri/src/commands/conversations.rs index acd0c643..965e995a 100644 --- a/src-tauri/src/commands/conversations.rs +++ b/src-tauri/src/commands/conversations.rs @@ -1434,34 +1434,51 @@ fn is_compressible_boundary_message(message: &Message) -> bool { && message.role != MessageRole::Tool } -fn last_compressible_message_id_before( +/// Pick `compressed_until_message_id` so the last `keep_last_n` compressible +/// messages (and everything from `force_retain_from_id` onward) stay cleartext. +/// +/// Returns `None` when there is nothing left to compress. +fn resolve_compressed_until_with_keep( db_messages: &[Message], start_index: usize, - before_message_id: &str, + keep_last_n: u32, + force_retain_from_id: Option<&str>, ) -> Option { - let end_index = db_messages - .iter() - .position(|message| message.id == before_message_id) - .unwrap_or(db_messages.len()); - db_messages - .iter() - .skip(start_index) - .take(end_index.saturating_sub(start_index)) - .filter(|message| is_compressible_boundary_message(message)) - .last() - .map(|message| message.id.clone()) -} + let force_idx = force_retain_from_id.and_then(|id| { + db_messages + .iter() + .position(|message| message.id == id) + }); -fn last_compressible_message_id_from_start( - db_messages: &[Message], - start_index: usize, -) -> Option { - db_messages + let compressible_indices: Vec = db_messages .iter() + .enumerate() .skip(start_index) - .filter(|message| is_compressible_boundary_message(message)) - .last() - .map(|message| message.id.clone()) + .filter(|(_, message)| is_compressible_boundary_message(message)) + .map(|(idx, _)| idx) + .collect(); + + if compressible_indices.is_empty() { + return None; + } + + let keep_start_by_n = compressible_indices + .len() + .saturating_sub(keep_last_n as usize); + let keep_start_by_force = force_idx + .and_then(|force| { + compressible_indices + .iter() + .position(|&idx| idx >= force) + }) + .unwrap_or(compressible_indices.len()); + let keep_start = keep_start_by_n.min(keep_start_by_force); + + if keep_start == 0 { + return None; + } + + Some(db_messages[compressible_indices[keep_start - 1]].id.clone()) } fn count_compressible_messages_from_start(db_messages: &[Message], start_index: usize) -> u32 { @@ -1724,28 +1741,13 @@ fn build_provider_context_messages_from_index( fn split_auto_compression_history( history_messages: &[ChatMessage], current_user_index: Option, + keep_last_n: u32, ) -> (Vec, Vec) { - let Some(current_index) = current_user_index else { - return (history_messages.to_vec(), Vec::new()); - }; - if current_index >= history_messages.len() { - return (history_messages.to_vec(), Vec::new()); - } - - let messages_to_compress = history_messages - .iter() - .enumerate() - .filter_map(|(index, message)| { - if index == current_index { - None - } else { - Some(message.clone()) - } - }) - .collect(); - let post_compression_history = vec![history_messages[current_index].clone()]; - - (messages_to_compress, post_compression_history) + crate::context_manager::split_history_keep_last( + history_messages, + keep_last_n, + current_user_index, + ) } #[tauri::command] @@ -4395,15 +4397,25 @@ pub async fn send_message( model_context_window, ) { - let (messages_to_compress, post_compression_history) = - split_auto_compression_history(&history_messages, current_user_history_index); - let compressed_until_message_id = last_compressible_message_id_before( + let keep_last_n = crate::context_manager::resolve_compression_keep_last_n( + conversation.compression_keep_last_n, + global_settings.default_compression_keep_last_n, + ); + let (messages_to_compress, post_compression_history) = split_auto_compression_history( + &history_messages, + current_user_history_index, + keep_last_n, + ); + let compressed_until_message_id = resolve_compressed_until_with_keep( &db_messages, context_boundary.start_index, - &user_message.id, + keep_last_n, + Some(&user_message.id), ); // Perform synchronous compression before sending - let compression_result = if messages_to_compress.is_empty() { + let compression_result = if messages_to_compress.is_empty() + || compressed_until_message_id.is_none() + { None } else { do_compress( @@ -4450,8 +4462,7 @@ pub async fn send_message( ); } - // After compression, history is now empty (marker splits it) - // Context = system + summary + current user message only + // Context = system + summary + retained trailing messages chat_messages = crate::context_manager::build_context( &chat_messages, &post_compression_history, @@ -5359,6 +5370,7 @@ async fn do_compress( messages_to_compress: history_messages.to_vec(), }; + let source_text = crate::context_manager::format_compression_source_text(&sum_req); let custom_prompt = settings.compression_prompt.as_deref(); let summary_messages = if let Some(prompt) = custom_prompt { crate::context_manager::build_summary_prompt_with_custom(&sum_req, prompt) @@ -5366,8 +5378,52 @@ async fn do_compress( crate::context_manager::build_summary_prompt(&sum_req) }; + let (response_content, comp_model_id) = run_compression_llm( + &comp_provider, + &comp_key, + &comp_key_id, + &comp_proxy, + &comp_model_id, + comp_use_max, + summary_messages, + settings, + ) + .await?; + + let token_count = aqbot_core::token_counter::estimate_tokens(&response_content); + let summary = aqbot_core::repo::conversation::upsert_summary( + db, + conversation_id, + &response_content, + compressed_until_message_id, + Some(token_count as u32), + Some(&comp_model_id), + Some(&source_text), + ) + .await + .map_err(|e| format!("Failed to save summary: {}", e))?; + + tracing::debug!( + "Compressed context for {} ({} tokens)", + conversation_id, + token_count + ); + Ok(summary) +} + +/// Shared LLM call for compression / retry. +async fn run_compression_llm( + comp_provider: &ProviderConfig, + comp_key: &str, + comp_key_id: &str, + comp_proxy: &Option, + comp_model_id: &str, + comp_use_max: Option, + summary_messages: Vec, + settings: &AppSettings, +) -> Result<(String, String), String> { let request = ChatRequest { - model: comp_model_id.clone(), + model: comp_model_id.to_string(), messages: summary_messages, stream: false, temperature: settings @@ -5386,8 +5442,8 @@ async fn do_compress( }; let ctx = ProviderRequestContext { - api_key: comp_key, - key_id: comp_key_id, + api_key: comp_key.to_string(), + key_id: comp_key_id.to_string(), provider_id: comp_provider.id.clone(), base_url: Some(resolve_base_url_for_type( &comp_provider.api_host, @@ -5395,7 +5451,7 @@ async fn do_compress( )), api_path: comp_provider.api_path.clone(), aws_region: comp_provider.aws_region.clone(), - proxy_config: comp_proxy, + proxy_config: comp_proxy.clone(), custom_headers: comp_provider .custom_headers .as_ref() @@ -5416,24 +5472,7 @@ async fn do_compress( return Err("Summary generation returned inline image data".to_string()); } - let token_count = aqbot_core::token_counter::estimate_tokens(&response.content); - let summary = aqbot_core::repo::conversation::upsert_summary( - db, - conversation_id, - &response.content, - compressed_until_message_id, - Some(token_count as u32), - Some(&comp_model_id), - ) - .await - .map_err(|e| format!("Failed to save summary: {}", e))?; - - tracing::debug!( - "Compressed context for {} ({} tokens)", - conversation_id, - token_count - ); - Ok(summary) + Ok((response.content, comp_model_id.to_string())) } /// Tauri command: manually compress the current conversation context. @@ -5483,45 +5522,48 @@ pub async fn compress_context( .ok() .flatten(); let context_boundary = resolve_context_boundary(&db_messages, existing_summary.as_ref()); + let keep_last_n = crate::context_manager::resolve_compression_keep_last_n( + conversation.compression_keep_last_n, + global_settings.default_compression_keep_last_n, + ); + let mut boundary_start_index = context_boundary.start_index; - let mut history_messages = build_provider_context_messages_from_index( + let mut compressed_until_message_id = resolve_compressed_until_with_keep( + &db_messages, + boundary_start_index, + keep_last_n, + None, + ); + + // Fall back: if nothing to compress after boundary (only keep-N left), try full history. + if compressed_until_message_id.is_none() && boundary_start_index > 0 { + boundary_start_index = 0; + compressed_until_message_id = + resolve_compressed_until_with_keep(&db_messages, 0, keep_last_n, None); + } + + let Some(compressed_until_message_id) = compressed_until_message_id else { + return Err("No messages to compress (not enough beyond keep-last-N)".to_string()); + }; + + let history_messages = build_provider_context_messages_from_index( &file_store, &db_messages, boundary_start_index, global_settings.document_attachment_reading_enabled, None, None, - None, + Some(&compressed_until_message_id), ) .map_err(|e| e.to_string())?; - // If nothing after the last boundary, try all messages. - if history_messages.is_empty() && boundary_start_index > 0 { - let all_without_markers = db_messages - .iter() - .filter(|message| !is_context_boundary_marker(message)) - .cloned() - .collect::>(); - boundary_start_index = 0; - history_messages = build_provider_context_messages( - &file_store, - &all_without_markers, - global_settings.document_attachment_reading_enabled, - None, - None, - None, - ) - .map_err(|e| e.to_string())?; - } - if history_messages.is_empty() { return Err("No messages to compress".to_string()); } - let compressed_until_message_id = - last_compressible_message_id_from_start(&db_messages, boundary_start_index); + let effective_existing_summary = existing_summary .as_ref() - .filter(|_| context_boundary.use_summary); + .filter(|_| context_boundary.use_summary && boundary_start_index == context_boundary.start_index); // Compress let use_max_completion_tokens = aqbot_core::repo::provider::get_model( @@ -5539,7 +5581,7 @@ pub async fn compress_context( &conversation_id, &history_messages, effective_existing_summary.map(|s| s.summary_text.as_str()), - compressed_until_message_id.as_deref(), + Some(&compressed_until_message_id), &provider, &decrypted_key, &key_row.id, @@ -5577,6 +5619,177 @@ pub async fn compress_context( Ok(summary) } +/// Tauri command: re-run compression on the stored source text with the current +/// global compression prompt. Does not change the boundary or insert a new marker. +#[tauri::command] +pub async fn retry_compression( + app: tauri::AppHandle, + state: State<'_, AppState>, + conversation_id: String, +) -> Result { + let conversation = + aqbot_core::repo::conversation::get_conversation(&state.sea_db, &conversation_id) + .await + .map_err(|e| e.to_string())?; + + let existing = aqbot_core::repo::conversation::get_summary(&state.sea_db, &conversation_id) + .await + .map_err(|e| e.to_string())? + .ok_or_else(|| "No compression summary to retry".to_string())?; + + let source_text = existing + .source_text + .as_deref() + .filter(|s| !s.trim().is_empty()) + .ok_or_else(|| { + "No compression source text saved (legacy summary). Compress again to enable retry." + .to_string() + })?; + + let provider = + aqbot_core::repo::provider::get_provider(&state.sea_db, &conversation.provider_id) + .await + .map_err(|e| e.to_string())?; + let fallback_key_id = provider + .keys + .first() + .map(|k| k.id.clone()) + .ok_or_else(|| "No API key configured".to_string())?; + let fallback_key_encrypted = provider + .keys + .first() + .map(|k| k.key_encrypted.clone()) + .ok_or_else(|| "No API key configured".to_string())?; + let decrypted_key = + aqbot_core::crypto::decrypt_key(&fallback_key_encrypted, &state.master_key) + .map_err(|e| e.to_string())?; + + let global_settings = aqbot_core::repo::settings::get_settings(&state.sea_db) + .await + .unwrap_or_default(); + let resolved_proxy = ProviderProxyConfig::resolve(&provider.proxy_config, &global_settings); + + // Resolve compression model (same cascade as do_compress) + let (comp_provider, comp_key, comp_key_id, comp_proxy, comp_model_id, comp_use_max) = if let ( + Some(ref pid), + Some(ref mid), + ) = ( + &global_settings.compression_provider_id, + &global_settings.compression_model_id, + ) { + match aqbot_core::repo::provider::get_provider(&state.sea_db, pid).await { + Ok(p) => { + let first_key = p.keys.first().cloned(); + match first_key { + Some(k) => { + let dk = aqbot_core::crypto::decrypt_key(&k.key_encrypted, &state.master_key) + .map_err(|e| e.to_string())?; + let override_umc = + aqbot_core::repo::provider::get_model(&state.sea_db, pid, mid) + .await + .ok() + .and_then(|m| m.param_overrides) + .and_then(|po| po.use_max_completion_tokens); + let proxy = ProviderProxyConfig::resolve(&p.proxy_config, &global_settings); + (p, dk, k.id, proxy, mid.clone(), override_umc) + } + None => ( + provider.clone(), + decrypted_key.clone(), + fallback_key_id.clone(), + resolved_proxy.clone(), + conversation.model_id.clone(), + aqbot_core::repo::provider::get_model( + &state.sea_db, + &conversation.provider_id, + &conversation.model_id, + ) + .await + .ok() + .and_then(|m| m.param_overrides) + .and_then(|po| po.use_max_completion_tokens), + ), + } + } + Err(_) => ( + provider.clone(), + decrypted_key.clone(), + fallback_key_id.clone(), + resolved_proxy.clone(), + conversation.model_id.clone(), + None, + ), + } + } else { + let use_max = aqbot_core::repo::provider::get_model( + &state.sea_db, + &conversation.provider_id, + &conversation.model_id, + ) + .await + .ok() + .and_then(|m| m.param_overrides) + .and_then(|po| po.use_max_completion_tokens); + ( + provider, + decrypted_key, + fallback_key_id, + resolved_proxy, + conversation.model_id.clone(), + use_max, + ) + }; + + let system_prompt = global_settings + .compression_prompt + .as_deref() + .filter(|s| !s.trim().is_empty()) + .unwrap_or( + "你是一个对话摘要助手。请将以下对话历史压缩为简洁摘要。\n\n\ + 要求:\n\ + 1. 保留所有用户明确表达的需求、偏好和决策\n\ + 2. 保留关键技术细节(代码片段、配置、错误信息等)\n\ + 3. 保留待办事项和未解决的问题\n\ + 4. 用简洁的要点形式组织\n\ + 5. 保持摘要简洁,不超过 500 字", + ); + + let summary_messages = + crate::context_manager::build_summary_prompt_from_source(source_text, system_prompt); + + let (response_content, used_model) = run_compression_llm( + &comp_provider, + &comp_key, + &comp_key_id, + &comp_proxy, + &comp_model_id, + comp_use_max, + summary_messages, + &global_settings, + ) + .await?; + + let token_count = aqbot_core::token_counter::estimate_tokens(&response_content); + let summary = aqbot_core::repo::conversation::update_summary_text( + &state.sea_db, + &conversation_id, + &response_content, + Some(token_count as u32), + Some(&used_model), + ) + .await + .map_err(|e| e.to_string())?; + + ensure_conversation_summary_safe_for_ipc(&summary)?; + + let _ = app.emit( + "conversation:summary-updated", + summary.clone(), + ); + + Ok(summary) +} + /// Tauri command: get the compression summary for a conversation. #[tauri::command] pub async fn get_compression_summary( @@ -5602,6 +5815,7 @@ fn ensure_conversation_summary_safe_for_ipc(summary: &ConversationSummary) -> Re "compressed_until_message_id", summary.compressed_until_message_id.as_deref(), ), + ("source_text", summary.source_text.as_deref()), ("model_used", summary.model_used.as_deref()), ] .into_iter() @@ -5832,6 +6046,7 @@ mod tests { is_archived: false, context_compression: false, context_message_limit: None, + compression_keep_last_n: None, category_id: None, parent_conversation_id: None, mode: "chat".to_string(), @@ -5991,6 +6206,7 @@ mod tests { conversation_id: "conv-1".to_string(), summary_text: "compressed old context".to_string(), compressed_until_message_id: boundary_message_id.map(str::to_string), + source_text: None, token_count: Some(12), model_used: Some("summary-model".to_string()), created_at: 1, @@ -7010,6 +7226,43 @@ mod tests { assert_eq!(serialized["content"], "final answer"); } + #[test] + fn resolve_compressed_until_with_keep_last_n() { + let messages = vec![ + test_message("u1", MessageRole::User, "u1", None, 0, true, None, None), + test_message("a1", MessageRole::Assistant, "a1", Some("u1"), 0, true, None, None), + test_message("u2", MessageRole::User, "u2", None, 0, true, None, None), + test_message("a2", MessageRole::Assistant, "a2", Some("u2"), 0, true, None, None), + test_message("u3", MessageRole::User, "u3", None, 0, true, None, None), + ]; + + // keep 0 → compress all + assert_eq!( + resolve_compressed_until_with_keep(&messages, 0, 0, None).as_deref(), + Some("u3") + ); + // keep 3 → compress u1,a1; until a1 + assert_eq!( + resolve_compressed_until_with_keep(&messages, 0, 3, None).as_deref(), + Some("a1") + ); + // keep 5 → nothing to compress + assert_eq!( + resolve_compressed_until_with_keep(&messages, 0, 5, None), + None + ); + // auto: force retain from u3, keep 0 → until a2 + assert_eq!( + resolve_compressed_until_with_keep(&messages, 0, 0, Some("u3")).as_deref(), + Some("a2") + ); + // auto: force u3, keep 3 → retain u2,a2,u3 → until a1 + assert_eq!( + resolve_compressed_until_with_keep(&messages, 0, 3, Some("u3")).as_deref(), + Some("a1") + ); + } + #[test] fn auto_compression_excludes_current_user_from_summary_and_keeps_it_for_request() { let history_messages = vec![ @@ -7036,8 +7289,9 @@ mod tests { }, ]; + // keep_last_n=0 still retains the current user turn let (messages_to_compress, post_compression_history) = - split_auto_compression_history(&history_messages, Some(2)); + split_auto_compression_history(&history_messages, Some(2), 0); assert_eq!(messages_to_compress.len(), 2); assert_eq!(post_compression_history.len(), 1); @@ -7051,6 +7305,12 @@ mod tests { ChatContent::Text(content) if content.contains("current user message") ) })); + + // keep_last_n=3 retains last 3 including current user + let (to_compress_n3, retained_n3) = + split_auto_compression_history(&history_messages, Some(2), 3); + assert!(to_compress_n3.is_empty()); + assert_eq!(retained_n3.len(), 3); } #[test] diff --git a/src-tauri/src/context_manager.rs b/src-tauri/src/context_manager.rs index 3753bdcb..1b297e7b 100644 --- a/src-tauri/src/context_manager.rs +++ b/src-tauri/src/context_manager.rs @@ -16,6 +16,27 @@ const THRESHOLD_RATIO: f64 = 0.70; /// Content string for the compression marker message. pub const COMPRESSION_MARKER: &str = ""; +/// Default number of trailing compressible messages to leave out of compression. +pub const DEFAULT_COMPRESSION_KEEP_LAST_N: u32 = 3; + +/// Resolve keep-last-N: +/// conversation override → global default → hardcoded 3. +/// Explicit `Some(0)` means keep none. +pub fn resolve_compression_keep_last_n( + conversation_value: Option, + global_default: Option, +) -> u32 { + conversation_value + .or(global_default) + .unwrap_or(DEFAULT_COMPRESSION_KEEP_LAST_N) +} + +/// Short instruction restated after the conversation body so models that +/// "continue the chat" instead of summarizing still see the constraint. +pub const COMPRESSION_FOOTER_REMINDER: &str = "\n\n---\n\ +请严格按系统指令执行:只输出对话摘要,不要继续回答对话内容中的问题,\ +不要扮演对话中的角色,不要输出摘要以外的任何内容。"; + /// Estimate the token count of a single `ChatMessage`. pub fn message_tokens(msg: &ChatMessage) -> usize { let text = match &msg.content { @@ -251,11 +272,49 @@ pub struct SummarizationRequest { pub messages_to_compress: Vec, } -/// Build the LLM prompt for generating a conversation summary. -pub fn build_summary_prompt(request: &SummarizationRequest) -> Vec { - let mut messages = Vec::new(); +/// Format the conversation body used as compression input (and stored as `source_text`). +pub fn format_compression_source_text(request: &SummarizationRequest) -> String { + let conversation_text: Vec = request + .messages_to_compress + .iter() + .map(format_message_for_summary) + .collect(); + + let mut parts = Vec::new(); + if let Some(ref summary) = request.existing_summary { + parts.push(format!("已有摘要:\n{}", summary)); + } + parts.push(format!( + "{}对话内容:\n{}", + if request.existing_summary.is_some() { + "新增" + } else { + "" + }, + conversation_text.join("\n") + )); + parts.join("\n\n") +} + +fn format_message_for_summary(m: &ChatMessage) -> String { + let content_text = match &m.content { + ChatContent::Text(s) => s.clone(), + ChatContent::Multipart(parts) => parts + .iter() + .filter_map(|p| p.text.as_deref()) + .collect::>() + .join(" "), + }; + let truncated = if content_text.len() > 2000 { + format!("{}...[已截断]", &content_text[..2000]) + } else { + content_text + }; + format!("{}: {}", m.role, truncated) +} - let instruction = if request.existing_summary.is_some() { +fn default_compression_instruction(has_existing_summary: bool) -> &'static str { + if has_existing_summary { "你是一个对话摘要助手。请将以下新增对话内容合并到已有摘要中。\n\n\ 要求:\n\ 1. 保留所有用户明确表达的需求、偏好和决策\n\ @@ -272,64 +331,15 @@ pub fn build_summary_prompt(request: &SummarizationRequest) -> Vec 3. 保留待办事项和未解决的问题\n\ 4. 用简洁的要点形式组织\n\ 5. 保持摘要简洁,不超过 500 字" - }; - - messages.push(ChatMessage { - role: "system".to_string(), - content: ChatContent::Text(instruction.to_string()), - reasoning_content: None, - tool_calls: None, - tool_call_id: None, - }); - - if let Some(ref summary) = request.existing_summary { - messages.push(ChatMessage { - role: "user".to_string(), - content: ChatContent::Text(format!("已有摘要:\n{}", summary)), - reasoning_content: None, - tool_calls: None, - tool_call_id: None, - }); } +} - let conversation_text: Vec = request - .messages_to_compress - .iter() - .map(|m| { - let content_text = match &m.content { - ChatContent::Text(s) => s.clone(), - ChatContent::Multipart(parts) => parts - .iter() - .filter_map(|p| p.text.as_deref()) - .collect::>() - .join(" "), - }; - let truncated = if content_text.len() > 2000 { - format!("{}...[已截断]", &content_text[..2000]) - } else { - content_text - }; - format!("{}: {}", m.role, truncated) - }) - .collect(); - - messages.push(ChatMessage { - role: "user".to_string(), - content: ChatContent::Text(format!( - "{}对话内容:\n{}", - if request.existing_summary.is_some() { - "新增" - } else { - "" - }, - conversation_text.join("\n") - )), - reasoning_content: None, - tool_calls: None, - tool_call_id: None, - }); - - messages +/// Build the LLM prompt for generating a conversation summary. +pub fn build_summary_prompt(request: &SummarizationRequest) -> Vec { + build_summary_prompt_with_system( + request, + default_compression_instruction(request.existing_summary.is_some()), + ) } /// Build summary prompt with a custom system instruction (from settings). @@ -337,64 +347,105 @@ pub fn build_summary_prompt_with_custom( request: &SummarizationRequest, custom_prompt: &str, ) -> Vec { - let mut messages = Vec::new(); + build_summary_prompt_with_system(request, custom_prompt) +} - messages.push(ChatMessage { - role: "system".to_string(), - content: ChatContent::Text(custom_prompt.to_string()), - reasoning_content: None, - tool_calls: None, - tool_call_id: None, - }); +/// Rebuild a compression prompt from stored `source_text` (retry path). +pub fn build_summary_prompt_from_source(source_text: &str, system_prompt: &str) -> Vec { + vec![ + ChatMessage { + role: "system".to_string(), + content: ChatContent::Text(system_prompt.to_string()), + reasoning_content: None, + tool_calls: None, + tool_call_id: None, + }, + ChatMessage { + role: "user".to_string(), + content: ChatContent::Text(format!( + "{}{}", + source_text, COMPRESSION_FOOTER_REMINDER + )), + reasoning_content: None, + tool_calls: None, + tool_call_id: None, + }, + ] +} - if let Some(ref summary) = request.existing_summary { - messages.push(ChatMessage { +fn build_summary_prompt_with_system( + request: &SummarizationRequest, + system_prompt: &str, +) -> Vec { + let source = format_compression_source_text(request); + vec![ + ChatMessage { + role: "system".to_string(), + content: ChatContent::Text(system_prompt.to_string()), + reasoning_content: None, + tool_calls: None, + tool_call_id: None, + }, + ChatMessage { role: "user".to_string(), - content: ChatContent::Text(format!("已有摘要:\n{}", summary)), + content: ChatContent::Text(format!("{}{}", source, COMPRESSION_FOOTER_REMINDER)), reasoning_content: None, tool_calls: None, tool_call_id: None, - }); + }, + ] +} + +/// Split provider history into (to_compress, retained) keeping the last +/// `keep_last_n` messages (group-aware via [`message_group_start`]). +/// +/// When `current_user_index` is set (auto path), the current user message and +/// everything after it is always retained, even if `keep_last_n` is 0. +pub fn split_history_keep_last( + history_messages: &[ChatMessage], + keep_last_n: u32, + current_user_index: Option, +) -> (Vec, Vec) { + if history_messages.is_empty() { + return (Vec::new(), Vec::new()); } - let conversation_text: Vec = request - .messages_to_compress - .iter() - .map(|m| { - let content_text = match &m.content { - ChatContent::Text(s) => s.clone(), - ChatContent::Multipart(parts) => parts - .iter() - .filter_map(|p| p.text.as_deref()) - .collect::>() - .join(" "), - }; - let truncated = if content_text.len() > 2000 { - format!("{}...[已截断]", &content_text[..2000]) - } else { - content_text - }; - format!("{}: {}", m.role, truncated) - }) - .collect(); + let from_current = current_user_index + .filter(|&idx| idx < history_messages.len()) + .map(|idx| history_messages.len() - idx) + .unwrap_or(0); + let retain_target = (keep_last_n as usize).max(from_current); - messages.push(ChatMessage { - role: "user".to_string(), - content: ChatContent::Text(format!( - "{}对话内容:\n{}", - if request.existing_summary.is_some() { - "新增" - } else { - "" - }, - conversation_text.join("\n") - )), - reasoning_content: None, - tool_calls: None, - tool_call_id: None, - }); - - messages + if retain_target == 0 { + return (history_messages.to_vec(), Vec::new()); + } + if retain_target >= history_messages.len() { + return (Vec::new(), history_messages.to_vec()); + } + + // Walk groups from the end until we have at least retain_target messages. + let mut total_msgs = 0usize; + let mut start_idx = history_messages.len(); + let mut end_idx = history_messages.len(); + + while end_idx > 0 { + let group_start = message_group_start(history_messages, end_idx - 1); + let group_len = end_idx - group_start; + if total_msgs > 0 && total_msgs + group_len > retain_target { + break; + } + total_msgs += group_len; + start_idx = group_start; + end_idx = group_start; + if total_msgs >= retain_target { + break; + } + } + + ( + history_messages[..start_idx].to_vec(), + history_messages[start_idx..].to_vec(), + ) } #[cfg(test)] @@ -451,6 +502,64 @@ mod tests { assert!(!should_auto_compress(&[], &history, Some(1_000_000))); } + #[test] + fn resolve_compression_keep_last_n_defaults_to_three() { + assert_eq!(resolve_compression_keep_last_n(None, None), 3); + assert_eq!(resolve_compression_keep_last_n(None, Some(5)), 5); + assert_eq!(resolve_compression_keep_last_n(Some(0), Some(5)), 0); + assert_eq!(resolve_compression_keep_last_n(Some(2), Some(5)), 2); + } + + #[test] + fn split_history_keep_last_retains_trailing_messages() { + let history = vec![ + text_message("user", "u1"), + text_message("assistant", "a1"), + text_message("user", "u2"), + text_message("assistant", "a2"), + text_message("user", "u3"), + ]; + + let (to_compress, retained) = split_history_keep_last(&history, 3, None); + assert_eq!(to_compress.len(), 2); + assert_eq!(retained.len(), 3); + match &retained[0].content { + ChatContent::Text(s) => assert_eq!(s, "u2"), + _ => panic!("expected text"), + } + + let (all, none) = split_history_keep_last(&history, 0, None); + assert_eq!(all.len(), 5); + assert!(none.is_empty()); + + // Auto path: keep_last_n=0 still retains current user at index 4 + let (compressed, post) = split_history_keep_last(&history, 0, Some(4)); + assert_eq!(compressed.len(), 4); + assert_eq!(post.len(), 1); + } + + #[test] + fn build_summary_prompt_appends_footer_reminder() { + let request = SummarizationRequest { + existing_summary: None, + messages_to_compress: vec![text_message("user", "hello")], + }; + let messages = build_summary_prompt(&request); + assert_eq!(messages.len(), 2); + match &messages[1].content { + ChatContent::Text(s) => { + assert!(s.contains("hello")); + assert!(s.contains(COMPRESSION_FOOTER_REMINDER.trim())); + } + _ => panic!("expected text"), + } + + let source = format_compression_source_text(&request); + assert!(source.contains("对话内容")); + assert!(source.contains("hello")); + assert!(!source.contains(COMPRESSION_FOOTER_REMINDER.trim())); + } + #[test] fn resolve_message_count_limit_prefers_conversation_over_global() { assert_eq!( diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 5ceec2a7..a2c88b66 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -304,6 +304,7 @@ pub fn run() { commands::conversations::send_system_message, commands::conversations::compress_context, commands::conversations::get_compression_summary, + commands::conversations::retry_compression, commands::conversations::get_context_usage, commands::conversations::delete_compression, commands::conversations::regenerate_conversation_title, diff --git a/src/components/chat/ChatView.tsx b/src/components/chat/ChatView.tsx index f67e7393..91555987 100644 --- a/src/components/chat/ChatView.tsx +++ b/src/components/chat/ChatView.tsx @@ -1,6 +1,6 @@ import React, { useMemo, useCallback, useRef, useState, useEffect, useSyncExternalStore } from 'react'; import { CloseCircleFilled, SyncOutlined } from '@ant-design/icons'; -import { Typography, Button, Dropdown, Input, App, Avatar, Alert, Popconfirm, Popover, theme, Tag, Image, Tooltip, Modal, Spin, Checkbox } from 'antd'; +import { Typography, Button, Dropdown, Input, App, Avatar, Alert, Popconfirm, Popover, theme, Tag, Image, Tooltip, Modal, Spin, Checkbox, Tabs } from 'antd'; import type { InputRef } from 'antd'; import { Pencil, Share2, FileImage, FileCode, FileText, FileType, Bot, Lightbulb, Code, Languages, Copy, Check, RotateCcw, User, Trash2, ChevronLeft, ChevronRight, ChevronDown, Scissors, Paperclip, AlertCircle, X, ArrowDown, ArrowUp, ArrowLeftRight, Zap, Sparkles, TextCursorInput, GitBranch, ChartNoAxesColumn, MessageSquare, ArrowUpRight, ArrowDownRight, Coins, Clock, Timer, Download, PanelLeftClose, PanelLeftOpen, ListChecks } from 'lucide-react'; import { ModelIcon } from '@lobehub/icons'; @@ -1965,12 +1965,16 @@ export function ChatView() { const branchConversation = useConversationStore((s) => s.branchConversation); const removeContextClear = useConversationStore((s) => s.removeContextClear); const getCompressionSummary = useConversationStore((s) => s.getCompressionSummary); + const retryCompression = useConversationStore((s) => s.retryCompression); const deleteCompression = useConversationStore((s) => s.deleteCompression); + const openCompressionSummaryToken = useConversationStore((s) => s.openCompressionSummaryToken); const listMessageVersionsBatch = useConversationStore((s) => s.listMessageVersionsBatch); const hydrateMessageVersions = useConversationStore((s) => s.hydrateMessageVersions); const [summaryModalOpen, setSummaryModalOpen] = useState(false); const [summaryModalText, setSummaryModalText] = useState(''); const [summaryModalSummary, setSummaryModalSummary] = useState(null); + const [summaryModalTab, setSummaryModalTab] = useState<'summary' | 'source'>('summary'); + const [summaryRetrying, setSummaryRetrying] = useState(false); const [previewPayload, setPreviewPayload] = useState(null); const [previewModalOpen, setPreviewModalOpen] = useState(false); const [mermaidPreviewSvg, setMermaidPreviewSvg] = useState(null); @@ -2185,6 +2189,31 @@ export function ChatView() { const titleInputRef = useRef(null); const skipTitleSaveRef = useRef(false); + const openCompressionSummaryModal = useCallback(async () => { + const convId = activeConversationId; + if (!convId) return; + const connectionGeneration = pageConnectionGenerationRef.current; + const summary = await getCompressionSummary(convId); + if ( + !pageActiveRef.current + || pageConnectionGenerationRef.current !== connectionGeneration + || useConversationStore.getState().activeConversationId !== convId + ) return; + setSummaryModalText(summary?.summary_text ?? t('chat.noSummary')); + setSummaryModalSummary(summary ?? null); + setSummaryModalTab('summary'); + setSummaryModalOpen(true); + }, [activeConversationId, getCompressionSummary, t]); + + // InputArea context circle requests opening this modal + const lastOpenCompressionSummaryTokenRef = useRef(0); + useEffect(() => { + if (openCompressionSummaryToken === 0) return; + if (openCompressionSummaryToken === lastOpenCompressionSummaryTokenRef.current) return; + lastOpenCompressionSummaryTokenRef.current = openCompressionSummaryToken; + void openCompressionSummaryModal(); + }, [openCompressionSummaryToken, openCompressionSummaryModal]); + usePageSuspendCleanup(() => { setStatsOpen(false); setStats(null); @@ -2196,6 +2225,8 @@ export function ChatView() { setCardBranchTitle(''); setSummaryModalOpen(false); setSummaryModalSummary(null); + setSummaryModalTab('summary'); + setSummaryRetrying(false); setPreviewModalOpen(false); setPreviewPayload(null); setMermaidPreviewOpen(false); @@ -4332,19 +4363,8 @@ export function ChatView() { > { - const convId = activeConversationId; - if (!convId) return; - const connectionGeneration = pageConnectionGenerationRef.current; - const summary = await getCompressionSummary(convId); - if ( - !pageActiveRef.current - || pageConnectionGenerationRef.current !== connectionGeneration - || useConversationStore.getState().activeConversationId !== convId - ) return; - setSummaryModalText(summary?.summary_text ?? t('chat.noSummary')); - setSummaryModalSummary(summary ?? null); - setSummaryModalOpen(true); + onClick={() => { + void openCompressionSummaryModal(); }} > {t('chat.contextCompressed')} @@ -4372,7 +4392,7 @@ export function ChatView() {
    ), }; - }, [activeConversationId, deleteCompression, getCompressionSummary, t, token.colorPrimary, token.colorPrimaryBorder, token.colorTextTertiary]); + }, [activeConversationId, deleteCompression, openCompressionSummaryModal, t, token.colorPrimary, token.colorPrimaryBorder, token.colorTextTertiary]); const contextCompressingRole = useCallback(() => { return { @@ -4862,8 +4882,48 @@ export function ChatView() { onCancel={() => { setSummaryModalOpen(false); setSummaryModalSummary(null); + setSummaryModalTab('summary'); }} - footer={null} + footer={ +
    + + +
    + } width={640} >
    @@ -4880,14 +4940,48 @@ export function ChatView() { )}
    )} - setSummaryModalTab(key as 'summary' | 'source')} + items={[ + { + key: 'summary', + label: t('chat.compressionSummary'), + children: ( + + ), + }, + { + key: 'source', + label: t('chat.compressionSource'), + children: summaryModalSummary?.source_text ? ( +
    +                    {summaryModalSummary.source_text}
    +                  
    + ) : ( +
    + {t('chat.noSourceText')} +
    + ), + }, + ]} />
    diff --git a/src/components/chat/ConversationSettingsModal.tsx b/src/components/chat/ConversationSettingsModal.tsx index 7393943b..f0e96f15 100644 --- a/src/components/chat/ConversationSettingsModal.tsx +++ b/src/components/chat/ConversationSettingsModal.tsx @@ -19,6 +19,9 @@ interface ConversationSettingsModalProps { const LEGACY_CONTEXT_LIMIT_KEY = (id: string) => `aqbot_context_limit_${id}`; /** Values ≥ this mean unlimited (matches backend CONTEXT_MESSAGE_LIMIT_UNLIMITED). */ const CONTEXT_LIMIT_UNLIMITED = 50; +/** Default keep-last-N when conversation field is null (matches backend DEFAULT_COMPRESSION_KEEP_LAST_N). */ +const DEFAULT_COMPRESSION_KEEP_LAST_N = 3; +const COMPRESSION_KEEP_LAST_N_MAX = 20; function resolveInitialContextLimit( conversationId: string, @@ -69,6 +72,7 @@ export function ConversationSettingsModal({ open, onClose }: ConversationSetting const [title, setTitle] = useState(''); const [systemPrompt, setSystemPrompt] = useState(''); const [contextLimit, setContextLimit] = useState(CONTEXT_LIMIT_UNLIMITED); + const [compressionKeepLastN, setCompressionKeepLastN] = useState(DEFAULT_COMPRESSION_KEEP_LAST_N); const [temperature, setTemperature] = useState(null); const [topP, setTopP] = useState(null); const [maxTokens, setMaxTokens] = useState(null); @@ -96,6 +100,15 @@ export function ConversationSettingsModal({ open, onClose }: ConversationSetting settings.default_context_count, ), ); + setCompressionKeepLastN( + conversation.compression_keep_last_n != null + && Number.isFinite(conversation.compression_keep_last_n) + ? Math.max(0, Math.min(COMPRESSION_KEEP_LAST_N_MAX, conversation.compression_keep_last_n)) + : settings.default_compression_keep_last_n != null + && Number.isFinite(settings.default_compression_keep_last_n) + ? Math.max(0, Math.min(COMPRESSION_KEEP_LAST_N_MAX, settings.default_compression_keep_last_n)) + : DEFAULT_COMPRESSION_KEEP_LAST_N, + ); // Load icon const iconStored = localStorage.getItem(CONV_ICON_KEY(conversation.id)); @@ -113,7 +126,7 @@ export function ConversationSettingsModal({ open, onClose }: ConversationSetting setIconValue(''); } } - }, [open, conversation, settings.default_context_count]); + }, [open, conversation, settings.default_context_count, settings.default_compression_keep_last_n]); if (!conversation) return null; @@ -136,6 +149,7 @@ export function ConversationSettingsModal({ open, onClose }: ConversationSetting top_p: topP, frequency_penalty: frequencyPenalty, context_message_limit: contextLimit, + compression_keep_last_n: compressionKeepLastN, }); // Drop legacy localStorage key after persisting to the database. try { @@ -284,6 +298,35 @@ export function ConversationSettingsModal({ open, onClose }: ConversationSetting
    + {/* Compression keep last N */} +
    +
    + {t('settings.compressionKeepLastN')} + + + + + {compressionKeepLastN} + +
    +
    + +
    +
    + {/* Temperature / Top P / Max Tokens / Frequency Penalty */} s.getContextUsage); + const requestOpenCompressionSummary = useConversationStore((s) => s.requestOpenCompressionSummary); const [serverContextUsage, setServerContextUsage] = useState<{ usedTokens: number; maxTokens: number; percent: number; + hasSummary: boolean; + messagesAfterBoundary: number; } | null>(null); const contextUsageRevision = `${messages.length}:${messages[messages.length - 1]?.id ?? ''}:${messages[messages.length - 1]?.status ?? ''}`; @@ -884,6 +887,8 @@ export function InputArea() { usedTokens: usage.used_tokens, maxTokens: usage.context_window, percent: Math.min(Math.round((usage.used_tokens / usage.context_window) * 100), 100), + hasSummary: usage.has_summary, + messagesAfterBoundary: usage.messages_after_boundary, }); }); }; @@ -929,7 +934,13 @@ export function InputArea() { } const percent = Math.min(Math.round((usedTokens / maxTokens) * 100), 100); - return { usedTokens, maxTokens, percent }; + return { + usedTokens, + maxTokens, + percent, + hasSummary: lastMarkerIdx !== -1 && activeMessages[lastMarkerIdx]?.content === '', + messagesAfterBoundary: effectiveMessages.length, + }; }, [messages, currentModel?.context_window, activeConversation?.system_prompt, serverContextUsage]); const { hasRealtimeVoice, hasReasoning, hasVision, hasFunctionCalling } = React.useMemo(() => ({ @@ -2048,12 +2059,49 @@ export function InputArea() { : contextTokenUsage.percent > 60 ? token.colorWarning : token.colorPrimary; + const canCompress = !!activeConversationId && !loading && !streaming && !compressing && messages.length > 0; return ( - {contextTokenUsage.usedTokens.toLocaleString()} / {contextTokenUsage.maxTokens.toLocaleString()} tokens ({contextTokenUsage.percent}%) - +
    +
    + {contextTokenUsage.usedTokens.toLocaleString()} / {contextTokenUsage.maxTokens.toLocaleString()} tokens ({contextTokenUsage.percent}%) +
    + {contextTokenUsage.hasSummary && ( +
    + {t('chat.hasSummary')} + {' · '} + {t('chat.messagesAfterBoundary', { count: contextTokenUsage.messagesAfterBoundary })} +
    + )} +
    + + +
    +
    } > = {}): Conversation { is_archived: false, context_compression: false, context_message_limit: null, + compression_keep_last_n: null, category_id: null, parent_conversation_id: null, mode: 'chat', diff --git a/src/components/chat/__tests__/conversationListModel.test.ts b/src/components/chat/__tests__/conversationListModel.test.ts index 2bc27114..1d7319c6 100644 --- a/src/components/chat/__tests__/conversationListModel.test.ts +++ b/src/components/chat/__tests__/conversationListModel.test.ts @@ -30,6 +30,7 @@ function conversation( is_archived: false, context_compression: false, context_message_limit: null, + compression_keep_last_n: null, category_id: null, parent_conversation_id: null, message_count: 0, diff --git a/src/components/settings/ConversationSettings.tsx b/src/components/settings/ConversationSettings.tsx index 6c0df818..b789d860 100644 --- a/src/components/settings/ConversationSettings.tsx +++ b/src/components/settings/ConversationSettings.tsx @@ -1,5 +1,5 @@ -import { Button, ColorPicker, Divider, Input, InputNumber, Switch, theme } from 'antd'; -import { FolderOpen, RotateCcw } from 'lucide-react'; +import { Button, ColorPicker, Divider, Input, InputNumber, Switch, Tooltip, theme } from 'antd'; +import { FolderOpen, Info, RotateCcw } from 'lucide-react'; import { useMemo } from 'react'; import { useTranslation } from 'react-i18next'; import { useSettingsStore } from '@/stores'; @@ -19,6 +19,17 @@ import { useSystemFonts } from '@/hooks/useSystemFonts'; import { SettingsGroup } from './SettingsGroup'; import { SettingsSelect } from './SettingsSelect'; +/** Matches backend DEFAULT_COMPRESSION_KEEP_LAST_N. */ +const DEFAULT_COMPRESSION_KEEP_LAST_N = 3; +const COMPRESSION_KEEP_LAST_N_MAX = 20; + +function normalizeCompressionKeepLastN(value: number | string | null | undefined) { + if (value == null || value === '') return DEFAULT_COMPRESSION_KEEP_LAST_N; + const numericValue = typeof value === 'number' ? value : Number(value); + if (!Number.isFinite(numericValue)) return DEFAULT_COMPRESSION_KEEP_LAST_N; + return Math.min(COMPRESSION_KEEP_LAST_N_MAX, Math.max(0, Math.floor(numericValue))); +} + const { TextArea } = Input; const CHAT_FONT_SIZE_MIN = 12; const CHAT_FONT_SIZE_MAX = 22; @@ -179,6 +190,33 @@ export function ConversationSettings() { /> + +
    + {t('settings.contextCompressionGroupDesc')} +
    +
    + + {t('settings.compressionKeepLastN')} + + + + + saveSettings({ + default_compression_keep_last_n: normalizeCompressionKeepLastN(value), + })} + style={{ width: 120 }} + /> +
    +
    + {t('settings.compressionKeepLastNHint')} +
    + +
    {t('settings.chatInputActionsScale')} diff --git a/src/components/settings/__tests__/ConversationSettings.test.tsx b/src/components/settings/__tests__/ConversationSettings.test.tsx index e9f5e72d..a2c0052a 100644 --- a/src/components/settings/__tests__/ConversationSettings.test.tsx +++ b/src/components/settings/__tests__/ConversationSettings.test.tsx @@ -53,6 +53,11 @@ vi.mock('react-i18next', () => ({ 'settings.codeFontFamily': '代码字体', 'settings.fontDefault': '系统默认', 'settings.groupMessageStyle': '消息样式', + 'settings.contextCompression': '上下文压缩', + 'settings.contextCompressionGroupDesc': '控制对话过长时如何压缩历史上下文。', + 'settings.compressionKeepLastN': '压缩时保留最近消息数', + 'settings.compressionKeepLastNTooltip': '保留最近 N 条不压入摘要', + 'settings.compressionKeepLastNHint': '此为全局默认', 'settings.agentSettings': 'Agent', 'settings.agentWorkspaceRoot': '默认工作目录', 'settings.agentWorkspaceRootDesc': '新 Agent 对话会在该目录下自动创建独立工作目录。留空时使用 ~/.aqbot/workspace。', @@ -186,6 +191,7 @@ vi.mock('antd', () => { ), Dropdown: ({ children }: { children?: React.ReactNode }) => <>{children}, + Tooltip: ({ children }: { children?: React.ReactNode; title?: React.ReactNode }) => <>{children}, theme: { useToken: () => ({ token: { @@ -275,9 +281,23 @@ describe('ConversationSettings', () => { agent_workspace_root: null, agent_workspace_name_strategy: 'uuid', agent_workspace_datetime_format: 'YYYY-MM-DD-HH-mm-ss', + default_compression_keep_last_n: null, }; }); + it('saves default compression keep-last-n from conversation settings', async () => { + render(); + + const input = screen.getByLabelText('压缩时保留最近消息数'); + fireEvent.change(input, { target: { value: '5' } }); + + await waitFor(() => { + expect(mocks.saveSettings).toHaveBeenCalledWith({ + default_compression_keep_last_n: 5, + }); + }); + }); + it('renders the additional features group below chat navigation', () => { render(); diff --git a/src/components/settings/__tests__/DefaultModelSettings.test.tsx b/src/components/settings/__tests__/DefaultModelSettings.test.tsx index 8d6f5427..f2c35640 100644 --- a/src/components/settings/__tests__/DefaultModelSettings.test.tsx +++ b/src/components/settings/__tests__/DefaultModelSettings.test.tsx @@ -58,6 +58,7 @@ describe('DefaultModelSettings', () => { compression_provider_id: null, compression_model_id: null, compression_prompt: null, + default_compression_keep_last_n: null, }; }); diff --git a/src/i18n/locales/en-US.json b/src/i18n/locales/en-US.json index b8032ada..edd48205 100644 --- a/src/i18n/locales/en-US.json +++ b/src/i18n/locales/en-US.json @@ -83,7 +83,17 @@ "compressing": "Compressing...", "contextMessages": "in context", "compressionSummary": "Compression Summary", + "compressionSource": "Compression source", + "retryCompression": "Retry compression", + "retryCompressionSuccess": "Summary regenerated with current prompt", + "retryCompressionFailed": "Retry compression failed", + "retryCompressionNoSource": "No source text (legacy summary). Compress again to enable retry.", + "viewCompressionSummary": "View summary", + "compressNow": "Compress now", + "hasSummary": "Has summary", + "messagesAfterBoundary": "{{count}} after boundary", "noSummary": "No summary available", + "noSourceText": "No source text", "exportMd": "Export Markdown", "exportJson": "Export JSON", "exportPng": "Export Image", @@ -1352,6 +1362,11 @@ "contextMessageLimit": "Context message count limit", "contextMessageLimitTooltip": "Limit how many messages are sent to the model (including the current user message). 0 = current only; 50 = unlimited.", "contextMessageLimitCurrentOnly": "Current only", + "contextCompression": "Context compression", + "contextCompressionGroupDesc": "Controls how long conversations compress history. Configure the compression model and prompt under Default Models.", + "compressionKeepLastN": "Messages to keep when compressing", + "compressionKeepLastNTooltip": "When compressing, leave the last N messages out of the summary so recent format and context stay intact. 0 compresses all; default is 3.", + "compressionKeepLastNHint": "This is the global default. You can override it in each conversation’s settings.", "iconGroupModel": "Model", "iconGroupProvider": "Provider", "indexStatus": { diff --git a/src/i18n/locales/zh-CN.json b/src/i18n/locales/zh-CN.json index e9929802..0254402b 100644 --- a/src/i18n/locales/zh-CN.json +++ b/src/i18n/locales/zh-CN.json @@ -83,7 +83,17 @@ "compressing": "正在压缩中...", "contextMessages": "条上下文", "compressionSummary": "压缩摘要", + "compressionSource": "压缩输入原文", + "retryCompression": "重新压缩", + "retryCompressionSuccess": "已用当前提示词重新生成摘要", + "retryCompressionFailed": "重新压缩失败", + "retryCompressionNoSource": "当前摘要无输入原文(旧数据),请重新手动压缩一次", + "viewCompressionSummary": "查看摘要", + "compressNow": "立即压缩", + "hasSummary": "已有摘要", + "messagesAfterBoundary": "边界后 {{count}} 条", "noSummary": "暂无摘要", + "noSourceText": "暂无输入原文", "exportMd": "导出 Markdown", "exportJson": "导出 JSON", "exportPng": "导出图片", @@ -1352,6 +1362,11 @@ "contextMessageLimit": "上下文的消息数量上限", "contextMessageLimitTooltip": "限制发送给模型的消息数量(含当前用户消息)。0 表示仅当前消息;设为 50 表示不限制。", "contextMessageLimitCurrentOnly": "仅当前", + "contextCompression": "上下文压缩", + "contextCompressionGroupDesc": "控制对话过长时如何压缩历史上下文。压缩所用模型与提示词请在「默认模型」中配置。", + "compressionKeepLastN": "压缩时保留最近消息数", + "compressionKeepLastNTooltip": "压缩时不把最近 N 条消息写入摘要,压缩后仍以原文形式保留在上下文中,有助于延续对话格式。0 表示全部压缩;默认 3。", + "compressionKeepLastNHint": "此为全局默认;可在单个对话的「对话设置」中单独覆盖。", "iconGroupModel": "模型", "iconGroupProvider": "厂商", "indexStatus": { diff --git a/src/lib/browserMock.ts b/src/lib/browserMock.ts index 9fb16505..1c9ad35a 100644 --- a/src/lib/browserMock.ts +++ b/src/lib/browserMock.ts @@ -977,6 +977,7 @@ export async function handleCommand(cmd: string, args?: Record Promise; /** Get the compression summary for a conversation */ getCompressionSummary: (conversationId: string) => Promise; + /** Re-run compression on stored source text with current global prompt */ + retryCompression: () => Promise; /** Get server-side context usage for a conversation */ getContextUsage: (conversationId: string) => Promise; /** Delete the compression summary and all marker messages */ deleteCompression: () => Promise; + /** Ask ChatView to open the compression summary modal for the active conversation */ + requestOpenCompressionSummary: () => void; ensureConversationsLoaded: (options?: EnsureLoadedOptions) => Promise; invalidateConversations: (reason: ResourceInvalidationReason) => void; fetchConversations: () => Promise; @@ -1641,6 +1647,7 @@ export const useConversationStore = create((set, get) => ({ newestLoadedMessageId: null, streaming: false, compressingConversationId: null, + openCompressionSummaryToken: 0, streamingMessageId: null, streamingConversationId: null, activeStreamId: null, @@ -1926,6 +1933,26 @@ export const useConversationStore = create((set, get) => ({ } }, + retryCompression: async () => { + const conversationId = get().activeConversationId; + if (!conversationId) return null; + if (get().loading) throw new Error('Conversation messages are still loading'); + set({ compressingConversationId: conversationId }); + try { + const summary = await invoke('retry_compression', { conversationId }); + set((state) => ({ + compressingConversationId: state.compressingConversationId === conversationId + ? null + : state.compressingConversationId, + })); + return summary; + } catch (e) { + set({ compressingConversationId: null }); + console.error('Failed to retry compression:', e); + throw e; + } + }, + getContextUsage: async (conversationId: string) => { try { return await invoke('get_context_usage', { conversationId }); @@ -1935,6 +1962,12 @@ export const useConversationStore = create((set, get) => ({ } }, + requestOpenCompressionSummary: () => { + set((state) => ({ + openCompressionSummaryToken: state.openCompressionSummaryToken + 1, + })); + }, + deleteCompression: async () => { const conversationId = get().activeConversationId; if (!conversationId) return; diff --git a/src/stores/settingsStore.ts b/src/stores/settingsStore.ts index 28436a5c..ca60f855 100644 --- a/src/stores/settingsStore.ts +++ b/src/stores/settingsStore.ts @@ -71,6 +71,7 @@ const DEFAULT_SETTINGS: AppSettings = { compression_top_p: null, compression_frequency_penalty: null, compression_prompt: null, + default_compression_keep_last_n: null, model_catalog_source: 'builtin', proxy_type: null, proxy_address: null, diff --git a/src/types/index.ts b/src/types/index.ts index fd0c056e..ac27bf42 100644 --- a/src/types/index.ts +++ b/src/types/index.ts @@ -386,6 +386,11 @@ export interface Conversation { context_compression: boolean; /** Per-conversation history message cap. null = use global default. ≥50 = unlimited. */ context_message_limit: number | null; + /** + * Keep the last N compressible messages out of compression. + * null = default (3). 0 = keep none. + */ + compression_keep_last_n: number | null; category_id: string | null; parent_conversation_id: string | null; mode?: 'chat' | 'agent' | 'role'; @@ -490,6 +495,8 @@ export interface ConversationSummary { conversation_id: string; summary_text: string; compressed_until_message_id: string | null; + /** Compression input text for viewing / retry. Absent on legacy summaries. */ + source_text?: string | null; token_count: number | null; model_used: string | null; created_at: number; @@ -532,6 +539,8 @@ export interface UpdateConversationInput { context_compression?: boolean; /** Set null to clear override and use global default. ≥50 = unlimited. */ context_message_limit?: number | null; + /** Set null to clear and use default keep-last-N (3). */ + compression_keep_last_n?: number | null; category_id?: string | null; mode?: 'chat' | 'agent' | 'role'; } @@ -711,6 +720,8 @@ export interface AppSettings { compression_top_p: number | null; compression_frequency_penalty: number | null; compression_prompt: string | null; + /** Global default keep-last-N when compressing. null → 3. Per-conversation override wins. */ + default_compression_keep_last_n: number | null; model_catalog_source: ModelCatalogSourcePreference; proxy_type: string | null; proxy_address: string | null; diff --git a/tests/performance/browserFixture.ts b/tests/performance/browserFixture.ts index 4360e759..ba62bbf9 100644 --- a/tests/performance/browserFixture.ts +++ b/tests/performance/browserFixture.ts @@ -58,6 +58,7 @@ function buildConversations( is_archived: false, context_compression: false, context_message_limit: null, + compression_keep_last_n: null, category_id: null, parent_conversation_id: null, mode: 'chat', From 8e52a7f8eeb5f8ca65b4c0d9a4ee68e171304a46 Mon Sep 17 00:00:00 2001 From: licoy Date: Fri, 7 Aug 2026 18:45:04 +0800 Subject: [PATCH 005/108] =?UTF-8?q?feat(gateway):=20=E6=94=AF=E6=8C=81?= =?UTF-8?q?=E5=90=8C=E5=90=8D=E6=A8=A1=E5=9E=8B=E8=81=9A=E5=90=88=E8=B7=AF?= =?UTF-8?q?=E7=94=B1=E4=B8=8E=E6=A8=A1=E5=9E=8B=E5=88=AB=E5=90=8D?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 说明: - 新增 gateway_auto_model_routing 开关:多服务商相同 model_id/别名时按 sort_order 选源,可重试错误时 failover - 模型支持多别名(aliases_json);网关请求别名会改写为真实 model_id 再转发上游 - /v1/models 在自动路由开启时对冲突名称聚合展示,并保留 provider/name 钉死入口 - 模型设置支持别名编辑;网关设置增加自动路由说明;统一模型设置表单控件默认尺寸 操作: - 升级后自动执行 migration(models.aliases_json) - 在服务商模型设置中配置别名;在网关设置中开启「自动模型路由」 Closes #137 --- src-tauri/crates/core/src/db.rs | 1 + src-tauri/crates/core/src/entity/models.rs | 2 + .../crates/core/src/repo/cherry_import.rs | 1 + .../crates/core/src/repo/kelivo_import.rs | 1 + src-tauri/crates/core/src/repo/provider.rs | 38 ++ .../crates/core/src/repo/provider_import.rs | 1 + src-tauri/crates/core/src/types.rs | 103 ++++ .../crates/core/tests/repo_integration.rs | 3 + src-tauri/crates/gateway/src/auto_route.rs | 293 ++++++++++ src-tauri/crates/gateway/src/handlers.rs | 502 ++++++++++++------ src-tauri/crates/gateway/src/lib.rs | 1 + src-tauri/crates/gateway/src/native.rs | 64 +-- src-tauri/crates/migration/src/lib.rs | 4 +- ...000001_add_model_aliases_and_auto_route.rs | 29 + src-tauri/crates/providers/src/anthropic.rs | 1 + .../crates/providers/src/bedrock/convert.rs | 1 + src-tauri/crates/providers/src/cohere.rs | 1 + src-tauri/crates/providers/src/gemini.rs | 1 + src-tauri/crates/providers/src/jina.rs | 1 + .../crates/providers/src/openai_compat.rs | 1 + .../crates/providers/src/openai_responses.rs | 1 + src-tauri/crates/providers/src/siliconflow.rs | 1 + src-tauri/crates/providers/src/voyage.rs | 1 + src-tauri/src/commands/conversations.rs | 1 + src-tauri/src/commands/drawing.rs | 3 + src-tauri/src/commands/providers.rs | 1 + src-tauri/src/model_catalog/inference.rs | 1 + src-tauri/src/model_catalog/tests/metadata.rs | 1 + src/components/common/ModelParamSliders.tsx | 2 - src/components/gateway/GatewaySettings.tsx | 18 + .../settings/ImageProtocolEditor.tsx | 2 +- src/components/settings/ProviderDetail.tsx | 99 +++- src/i18n/locales/en-US.json | 10 +- src/i18n/locales/zh-CN.json | 10 +- src/lib/browserMock.ts | 3 + src/stores/settingsStore.ts | 1 + src/types/index.ts | 4 + website/docs/guide/gateway.md | 25 + website/docs/zh/guide/gateway.md | 25 + 39 files changed, 1049 insertions(+), 209 deletions(-) create mode 100644 src-tauri/crates/gateway/src/auto_route.rs create mode 100644 src-tauri/crates/migration/src/m20260809_000001_add_model_aliases_and_auto_route.rs diff --git a/src-tauri/crates/core/src/db.rs b/src-tauri/crates/core/src/db.rs index 8a106907..3732f844 100644 --- a/src-tauri/crates/core/src/db.rs +++ b/src-tauri/crates/core/src/db.rs @@ -135,6 +135,7 @@ impl BuiltinModel { param_overrides: self.param_overrides.clone(), image_config: None, metadata_state: None, + aliases: Vec::new(), } } } diff --git a/src-tauri/crates/core/src/entity/models.rs b/src-tauri/crates/core/src/entity/models.rs index 584cc026..1db7e923 100644 --- a/src-tauri/crates/core/src/entity/models.rs +++ b/src-tauri/crates/core/src/entity/models.rs @@ -18,6 +18,8 @@ pub struct Model { pub param_overrides: Option, pub image_config_json: Option, pub metadata_state_json: Option, + /// JSON array of gateway request aliases for this model. + pub aliases_json: Option, } #[derive(Copy, Clone, Debug, EnumIter, DeriveRelation)] diff --git a/src-tauri/crates/core/src/repo/cherry_import.rs b/src-tauri/crates/core/src/repo/cherry_import.rs index 4a2c360d..49171ce3 100644 --- a/src-tauri/crates/core/src/repo/cherry_import.rs +++ b/src-tauri/crates/core/src/repo/cherry_import.rs @@ -1536,6 +1536,7 @@ where .and_then(|value| serde_json::to_string(&value).ok())), image_config_json: Set(None), metadata_state_json: Set(None), + aliases_json: Set(None), }) .on_conflict( OnConflict::columns([models::Column::ProviderId, models::Column::ModelId]) diff --git a/src-tauri/crates/core/src/repo/kelivo_import.rs b/src-tauri/crates/core/src/repo/kelivo_import.rs index 1c7dfe0d..b750c2a1 100644 --- a/src-tauri/crates/core/src/repo/kelivo_import.rs +++ b/src-tauri/crates/core/src/repo/kelivo_import.rs @@ -1147,6 +1147,7 @@ where .and_then(|value| serde_json::to_string(&value).ok())), image_config_json: Set(None), metadata_state_json: Set(None), + aliases_json: Set(None), }) .on_conflict( OnConflict::columns([models::Column::ProviderId, models::Column::ModelId]) diff --git a/src-tauri/crates/core/src/repo/provider.rs b/src-tauri/crates/core/src/repo/provider.rs index 1ea1ea43..9caad299 100644 --- a/src-tauri/crates/core/src/repo/provider.rs +++ b/src-tauri/crates/core/src/repo/provider.rs @@ -71,6 +71,12 @@ fn model_from_entity(m: models::Model) -> Model { }) .ok() }); + let aliases = m + .aliases_json + .as_deref() + .and_then(|s| serde_json::from_str::>(s).ok()) + .map(normalize_model_aliases) + .unwrap_or_default(); Model { provider_id: m.provider_id, model_id: m.model_id, @@ -88,6 +94,7 @@ fn model_from_entity(m: models::Model) -> Model { .image_config_json .and_then(|value| serde_json::from_str(&value).ok()), metadata_state, + aliases, } } @@ -666,6 +673,12 @@ where .metadata_state .as_ref() .and_then(|state| serde_json::to_string(state).ok()); + let aliases = normalize_model_aliases(model.aliases.iter().map(|s| s.as_str())); + let aliases_json = if aliases.is_empty() { + None + } else { + Some(serde_json::to_string(&aliases).unwrap_or_else(|_| "[]".to_string())) + }; models::ActiveModel { provider_id: Set(provider_id.to_string()), @@ -680,6 +693,7 @@ where param_overrides: Set(param_overrides), image_config_json: Set(image_config_json), metadata_state_json: Set(metadata_state_json), + aliases_json: Set(aliases_json), } .insert(conn) .await?; @@ -688,11 +702,31 @@ where Ok(()) } +fn validate_provider_model_aliases(input_models: &[Model]) -> Result<()> { + let siblings: Vec<(String, Vec)> = input_models + .iter() + .map(|m| { + ( + m.model_id.clone(), + normalize_model_aliases(m.aliases.iter().map(|s| s.as_str())), + ) + }) + .collect(); + for model in input_models { + let aliases = normalize_model_aliases(model.aliases.iter().map(|s| s.as_str())); + validate_model_aliases(&model.model_id, &aliases, &siblings) + .map_err(AQBotError::Validation)?; + } + Ok(()) +} + pub async fn save_models( db: &DatabaseConnection, provider_id: &str, input_models: &[Model], ) -> Result<()> { + validate_provider_model_aliases(input_models)?; + let provider_id = provider_id.to_string(); let input_models = input_models.to_vec(); @@ -721,6 +755,8 @@ pub async fn save_models_from_user_selection( provider_id: &str, input_models: &[Model], ) -> Result<()> { + validate_provider_model_aliases(input_models)?; + let provider = get_provider(db, provider_id).await?; let Some(builtin_id) = provider.builtin_id.clone() else { return save_models(db, provider_id, input_models).await; @@ -1192,6 +1228,7 @@ mod tests { max_output_tokens: ModelMetadataSource::Catalog, ..ModelMetadataState::default() }), + aliases: Vec::new(), }], ) .await @@ -1261,6 +1298,7 @@ mod tests { max_output_tokens: ModelMetadataSource::Catalog, ..ModelMetadataState::default() }), + aliases: Vec::new(), }], ) .await diff --git a/src-tauri/crates/core/src/repo/provider_import.rs b/src-tauri/crates/core/src/repo/provider_import.rs index df68009b..2867807c 100644 --- a/src-tauri/crates/core/src/repo/provider_import.rs +++ b/src-tauri/crates/core/src/repo/provider_import.rs @@ -226,6 +226,7 @@ async fn merge_candidate_models( param_overrides: empty_param_overrides_for_import(&provider_config.provider_type), image_config: None, metadata_state: None, + aliases: Vec::new(), }); } diff --git a/src-tauri/crates/core/src/types.rs b/src-tauri/crates/core/src/types.rs index 29358050..7fe7dc6d 100644 --- a/src-tauri/crates/core/src/types.rs +++ b/src-tauri/crates/core/src/types.rs @@ -256,6 +256,104 @@ pub struct Model { /// until the user explicitly restores automatic detection. #[serde(default)] pub metadata_state: Option, + /// Gateway request aliases. Clients may send an alias as `model`; the gateway + /// rewrites the upstream request to the real `model_id`. Empty by default. + #[serde(default)] + pub aliases: Vec, +} + +/// Maximum length of a single model alias. +pub const MODEL_ALIAS_MAX_LEN: usize = 128; + +/// Normalize aliases: trim, drop empty, de-duplicate while preserving order. +pub fn normalize_model_aliases(aliases: impl IntoIterator>) -> Vec { + let mut seen = std::collections::HashSet::new(); + let mut out = Vec::new(); + for alias in aliases { + let trimmed = alias.as_ref().trim(); + if trimmed.is_empty() { + continue; + } + if seen.insert(trimmed.to_string()) { + out.push(trimmed.to_string()); + } + } + out +} + +/// Validate aliases for a model within one provider's model list. +/// +/// Rules: +/// - each alias non-empty after trim, length ≤ [`MODEL_ALIAS_MAX_LEN`] +/// - alias ≠ own `model_id` +/// - unique among other models' `model_id` and aliases on the same provider +pub fn validate_model_aliases( + model_id: &str, + aliases: &[String], + sibling_models: &[(String, Vec)], +) -> Result<(), String> { + let normalized = normalize_model_aliases(aliases.iter().map(|s| s.as_str())); + for alias in &normalized { + if alias.len() > MODEL_ALIAS_MAX_LEN { + return Err(format!( + "Alias '{}' exceeds max length of {}", + alias, MODEL_ALIAS_MAX_LEN + )); + } + if alias == model_id { + return Err(format!( + "Alias '{}' must not equal the model's own model_id", + alias + )); + } + for (other_id, other_aliases) in sibling_models { + if other_id == model_id { + continue; + } + if other_id == alias { + return Err(format!( + "Alias '{}' conflicts with another model's model_id on this provider", + alias + )); + } + if other_aliases.iter().any(|a| a == alias) { + return Err(format!( + "Alias '{}' is already used by another model on this provider", + alias + )); + } + } + } + Ok(()) +} + +/// Whether a model is addressable by the given request name (real id or alias). +pub fn model_matches_request_name(model: &Model, name: &str) -> bool { + model.model_id == name || model.aliases.iter().any(|a| a == name) +} + +#[cfg(test)] +mod model_alias_tests { + use super::*; + + #[test] + fn normalize_trims_and_dedups() { + assert_eq!( + normalize_model_aliases([" a ", "a", "", "b"]), + vec!["a".to_string(), "b".to_string()] + ); + } + + #[test] + fn validate_rejects_own_model_id_and_sibling_conflict() { + let siblings = vec![ + ("gpt-5.5".to_string(), vec!["5.5".to_string()]), + ("other".to_string(), vec![]), + ]; + assert!(validate_model_aliases("gpt-5.5", &["gpt-5.5".into()], &siblings).is_err()); + assert!(validate_model_aliases("other", &["5.5".into()], &siblings).is_err()); + assert!(validate_model_aliases("other", &["fast".into()], &siblings).is_ok()); + } } #[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] @@ -1642,6 +1740,10 @@ pub struct AppSettings { pub gateway_ssl_key_path: Option, pub gateway_ssl_port: u16, pub gateway_force_ssl: bool, + /// When true, the gateway pools providers that share the same model id or + /// alias and fails over on retriable upstream errors. + #[serde(default)] + pub gateway_auto_model_routing: bool, pub always_on_top: bool, pub tray_enabled: bool, pub global_shortcuts_enabled: bool, @@ -1803,6 +1905,7 @@ impl Default for AppSettings { gateway_ssl_key_path: None, gateway_ssl_port: 8443, gateway_force_ssl: false, + gateway_auto_model_routing: false, always_on_top: false, tray_enabled: true, global_shortcuts_enabled: true, diff --git a/src-tauri/crates/core/tests/repo_integration.rs b/src-tauri/crates/core/tests/repo_integration.rs index c83fb25b..c369c7fb 100644 --- a/src-tauri/crates/core/tests/repo_integration.rs +++ b/src-tauri/crates/core/tests/repo_integration.rs @@ -229,9 +229,12 @@ async fn test_provider_model_operations() { model_type: ModelType::Chat, capabilities: vec![ModelCapability::TextChat], context_window: Some(4096), + max_output_tokens: None, enabled: true, param_overrides: None, image_config: None, + metadata_state: None, + aliases: Vec::new(), }]; // save models diff --git a/src-tauri/crates/gateway/src/auto_route.rs b/src-tauri/crates/gateway/src/auto_route.rs new file mode 100644 index 00000000..dfecafe4 --- /dev/null +++ b/src-tauri/crates/gateway/src/auto_route.rs @@ -0,0 +1,293 @@ +//! Automatic multi-provider model routing for the gateway. +//! +//! When the same model id or alias is configured on multiple providers and +//! `gateway_auto_model_routing` is enabled, requests form a candidate pool ordered +//! by `provider.sort_order`. Retriable upstream failures mark a candidate as +//! cooled-down and try the next one. + +use aqbot_core::types::{model_matches_request_name, Model, ProviderConfig}; +use std::collections::HashMap; +use std::sync::{Mutex, OnceLock}; +use std::time::{Duration, Instant}; + +/// Default cooldown after a retriable upstream failure. +pub const DEFAULT_COOLDOWN: Duration = Duration::from_secs(45); + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct RouteCandidate { + pub provider_id: String, + pub provider_name: String, + pub sort_order: i32, + pub real_model_id: String, + /// The client-facing name that matched (model_id or alias). + pub request_name: String, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum MatchVia { + ModelId, + Alias, +} + +/// In-process circuit breaker: `(provider_id, real_model_id) -> reopen_at`. +fn cooldown_map() -> &'static Mutex> { + static MAP: OnceLock>> = OnceLock::new(); + MAP.get_or_init(|| Mutex::new(HashMap::new())) +} + +pub fn mark_failure(provider_id: &str, model_id: &str) { + mark_failure_for(provider_id, model_id, DEFAULT_COOLDOWN); +} + +pub fn mark_failure_for(provider_id: &str, model_id: &str, cooldown: Duration) { + if let Ok(mut map) = cooldown_map().lock() { + map.insert( + (provider_id.to_string(), model_id.to_string()), + Instant::now() + cooldown, + ); + } +} + +pub fn mark_success(provider_id: &str, model_id: &str) { + if let Ok(mut map) = cooldown_map().lock() { + map.remove(&(provider_id.to_string(), model_id.to_string())); + } +} + +pub fn is_cooled_down(provider_id: &str, model_id: &str) -> bool { + let Ok(mut map) = cooldown_map().lock() else { + return false; + }; + let key = (provider_id.to_string(), model_id.to_string()); + match map.get(&key) { + Some(until) if *until > Instant::now() => true, + Some(_) => { + map.remove(&key); + false + } + None => false, + } +} + +/// Clear all cooldowns (tests only). +#[cfg(test)] +pub fn clear_cooldowns() { + if let Ok(mut map) = cooldown_map().lock() { + map.clear(); + } +} + +/// Find the real model_id on a model for a request name, if any. +pub fn resolve_real_model_id(model: &Model, name: &str) -> Option { + if model.model_id == name { + return Some(model.model_id.clone()); + } + if model.aliases.iter().any(|a| a == name) { + return Some(model.model_id.clone()); + } + None +} + +/// Collect enabled models matching `request_name` (real id or alias) across providers. +pub fn collect_candidates( + providers: &[ProviderConfig], + request_name: &str, +) -> Vec { + let mut out = Vec::new(); + for provider in providers.iter().filter(|p| p.enabled) { + for model in provider.models.iter().filter(|m| m.enabled) { + if !model_matches_request_name(model, request_name) { + continue; + } + out.push(RouteCandidate { + provider_id: provider.id.clone(), + provider_name: provider.name.clone(), + sort_order: provider.sort_order, + real_model_id: model.model_id.clone(), + request_name: request_name.to_string(), + }); + } + } + order_candidates(&mut out); + out +} + +pub fn order_candidates(cands: &mut [RouteCandidate]) { + cands.sort_by(|a, b| { + a.sort_order + .cmp(&b.sort_order) + .then_with(|| a.provider_id.cmp(&b.provider_id)) + .then_with(|| a.real_model_id.cmp(&b.real_model_id)) + }); +} + +/// Order candidates, skipping cooled-down ones when alternatives remain. +/// If every candidate is cooled down, return them all so we still try. +pub fn candidates_for_attempt(cands: &[RouteCandidate]) -> Vec { + let healthy: Vec = cands + .iter() + .filter(|c| !is_cooled_down(&c.provider_id, &c.real_model_id)) + .cloned() + .collect(); + if healthy.is_empty() { + cands.to_vec() + } else { + healthy + } +} + +/// Heuristic: whether an upstream/provider error is worth retrying on another source. +pub fn is_retriable_error_message(message: &str) -> bool { + let lower = message.to_lowercase(); + const NEEDLES: &[&str] = &[ + "429", + "408", + "500", + "502", + "503", + "504", + "timeout", + "timed out", + "connection", + "connect", + "reset", + "refused", + "unavailable", + "rate limit", + "rate_limit", + "overloaded", + "temporarily", + "bad gateway", + "gateway timeout", + "no active api key", + "no active key", + ]; + NEEDLES.iter().any(|n| lower.contains(n)) +} + +/// Distinct request names (model ids + aliases) exposed by enabled models. +/// Returns `(name, count_of_providers_exposing_it)`. +pub fn request_name_provider_counts(providers: &[ProviderConfig]) -> HashMap { + let mut counts: HashMap = HashMap::new(); + for provider in providers.iter().filter(|p| p.enabled) { + let mut names_on_provider = std::collections::HashSet::new(); + for model in provider.models.iter().filter(|m| m.enabled) { + names_on_provider.insert(model.model_id.clone()); + for alias in &model.aliases { + names_on_provider.insert(alias.clone()); + } + } + for name in names_on_provider { + *counts.entry(name).or_default() += 1; + } + } + counts +} + +#[cfg(test)] +mod tests { + use super::*; + use aqbot_core::types::{Model, ModelType, ProviderConfig, ProviderType}; + + fn provider(id: &str, sort: i32, models: Vec<(&str, &[&str])>) -> ProviderConfig { + ProviderConfig { + id: id.into(), + name: id.into(), + provider_type: ProviderType::OpenAI, + api_host: "https://example.com".into(), + api_path: None, + aws_region: None, + enabled: true, + models: models + .into_iter() + .map(|(mid, aliases)| Model { + provider_id: id.into(), + model_id: mid.into(), + name: mid.into(), + group_name: None, + model_type: ModelType::Chat, + capabilities: vec![], + context_window: None, + max_output_tokens: None, + enabled: true, + param_overrides: None, + image_config: None, + metadata_state: None, + aliases: aliases.iter().map(|s| (*s).to_string()).collect(), + }) + .collect(), + keys: vec![], + proxy_config: None, + custom_headers: None, + icon: None, + builtin_id: None, + sort_order: sort, + created_at: 0, + updated_at: 0, + } + } + + #[test] + fn collect_matches_model_id_and_alias() { + let providers = vec![ + provider("a", 0, vec![("gpt-5.5", &["5.5"])]), + provider("b", 1, vec![("gpt-5.5-turbo", &["5.5"])]), + ]; + let by_id = collect_candidates(&providers, "gpt-5.5"); + assert_eq!(by_id.len(), 1); + assert_eq!(by_id[0].provider_id, "a"); + assert_eq!(by_id[0].real_model_id, "gpt-5.5"); + + let by_alias = collect_candidates(&providers, "5.5"); + assert_eq!(by_alias.len(), 2); + assert_eq!(by_alias[0].provider_id, "a"); + assert_eq!(by_alias[1].provider_id, "b"); + assert_eq!(by_alias[1].real_model_id, "gpt-5.5-turbo"); + } + + #[test] + fn order_uses_sort_order() { + let providers = vec![ + provider("b", 10, vec![("m", &[])]), + provider("a", 0, vec![("m", &[])]), + ]; + let c = collect_candidates(&providers, "m"); + assert_eq!(c[0].provider_id, "a"); + assert_eq!(c[1].provider_id, "b"); + } + + #[test] + fn cooldown_skips_unhealthy_when_alternatives_exist() { + clear_cooldowns(); + let providers = vec![ + provider("a", 0, vec![("m", &[])]), + provider("b", 1, vec![("m", &[])]), + ]; + let all = collect_candidates(&providers, "m"); + mark_failure_for("a", "m", Duration::from_secs(60)); + let attempt = candidates_for_attempt(&all); + assert_eq!(attempt.len(), 1); + assert_eq!(attempt[0].provider_id, "b"); + clear_cooldowns(); + } + + #[test] + fn retriable_error_detection() { + assert!(is_retriable_error_message("HTTP 429 rate limit")); + assert!(is_retriable_error_message("connection refused")); + assert!(is_retriable_error_message("request timed out")); + assert!(!is_retriable_error_message("invalid_request: bad prompt")); + assert!(!is_retriable_error_message("401 unauthorized")); + } + + #[test] + fn request_name_counts_include_aliases() { + let providers = vec![ + provider("a", 0, vec![("gpt-5.5", &["5.5"])]), + provider("b", 1, vec![("other", &["5.5"])]), + ]; + let counts = request_name_provider_counts(&providers); + assert_eq!(counts.get("5.5"), Some(&2)); + assert_eq!(counts.get("gpt-5.5"), Some(&1)); + } +} diff --git a/src-tauri/crates/gateway/src/handlers.rs b/src-tauri/crates/gateway/src/handlers.rs index f6315957..ea0fa391 100644 --- a/src-tauri/crates/gateway/src/handlers.rs +++ b/src-tauri/crates/gateway/src/handlers.rs @@ -27,14 +27,13 @@ pub async fn health_check() -> impl IntoResponse { /// GET /v1/models — list enabled models from all enabled providers. /// -/// Model IDs are emitted as plain `model_id` when globally unique across all -/// enabled providers, or as `provider_slug/model_id` when the same `model_id` -/// exists on more than one enabled provider (collision). The legacy -/// `provider_uuid:model_id` format is **no longer emitted**. +/// Model IDs and aliases are listed. When the same request name collides across +/// providers: +/// - **auto routing off**: emit `public_id/name` (existing collision rule). +/// - **auto routing on**: emit a single bare name with `owned_by: "aqbot"`, and +/// still emit namespaced `public_id/name` entries so clients can pin a source. /// -/// Results are sorted deterministically: primary key is the displayed model ID -/// (lexicographic), secondary key is the provider name (tiebreaker for the rare -/// case of identical display IDs across multiple providers). +/// Results are sorted deterministically by displayed model ID, then owner. pub async fn list_models(State(state): State) -> impl IntoResponse { let providers = match aqbot_core::repo::provider::list_providers(&state.db).await { Ok(p) => p, @@ -47,26 +46,80 @@ pub async fn list_models(State(state): State) -> impl IntoRespo } }; - let display_map = build_model_display_map(&providers); + let auto_routing = aqbot_core::repo::settings::get_settings(&state.db) + .await + .map(|s| s.gateway_auto_model_routing) + .unwrap_or(false); + + let models = build_gateway_model_list(&providers, auto_routing); + + Json(json!({ + "object": "list", + "data": models, + })) + .into_response() +} + +/// Build OpenAI-style model list entries for the gateway. +pub(crate) fn build_gateway_model_list( + providers: &[ProviderConfig], + auto_routing: bool, +) -> Vec { + use crate::auto_route::request_name_provider_counts; + + let public_id_map = build_provider_public_id_map(providers); + let name_counts = request_name_provider_counts(providers); let mut models: Vec = Vec::new(); + let mut emitted_aggregated: HashSet = HashSet::new(); + for provider in providers.iter().filter(|p| p.enabled) { + let public_id = public_id_map + .get(&provider.id) + .cloned() + .unwrap_or_else(|| provider.name.clone()); + for model in provider.models.iter().filter(|m| m.enabled) { - let key = (provider.id.clone(), model.model_id.clone()); - let display_id = display_map - .get(&key) - .cloned() - .unwrap_or_else(|| model.model_id.clone()); - models.push(json!({ - "id": display_id, - "object": "model", - "created": provider.created_at, - "owned_by": provider.name, - })); + let mut names = vec![model.model_id.clone()]; + names.extend(model.aliases.iter().cloned()); + + for name in names { + let count = *name_counts.get(&name).unwrap_or(&0); + if auto_routing && count > 1 { + if emitted_aggregated.insert(name.clone()) { + models.push(json!({ + "id": name, + "object": "model", + "created": provider.created_at, + "owned_by": "aqbot", + })); + } + // Always keep namespaced pin entry when auto-routing aggregates. + models.push(json!({ + "id": format!("{}/{}", public_id, name), + "object": "model", + "created": provider.created_at, + "owned_by": provider.name, + })); + } else if count > 1 { + models.push(json!({ + "id": format!("{}/{}", public_id, name), + "object": "model", + "created": provider.created_at, + "owned_by": provider.name, + })); + } else { + models.push(json!({ + "id": name, + "object": "model", + "created": provider.created_at, + "owned_by": provider.name, + })); + } + } } } - // Deterministic ordering: display ID first, provider name as tiebreaker. models.sort_by(|a, b| { let id_a = a["id"].as_str().unwrap_or(""); let id_b = b["id"].as_str().unwrap_or(""); @@ -74,12 +127,7 @@ pub async fn list_models(State(state): State) -> impl IntoRespo let ob_b = b["owned_by"].as_str().unwrap_or(""); id_a.cmp(id_b).then(ob_a.cmp(ob_b)) }); - - Json(json!({ - "object": "list", - "data": models, - })) - .into_response() + models } /// POST /v1/chat/completions — main proxy handler @@ -117,97 +165,129 @@ pub async fn chat_completions( let known_public_ids: HashSet = public_id_map.values().cloned().collect(); // Parse model field: supports "provider_public_id/model_id" (preferred), - // legacy "provider_id:model_id" (compat), or bare "model_id". + // or bare "model_id" / alias. let parsed = parse_model_field(&request.model, &known_public_ids); - // Resolve the provider and canonical model_id. - let (provider, model_id) = match resolve_provider_for_model(&providers, &public_id_map, &parsed) - { - Ok(pair) => pair, - Err(resp) => return resp, - }; + let global_settings = aqbot_core::repo::settings::get_settings(&state.db) + .await + .unwrap_or_default(); + let auto_routing = global_settings.gateway_auto_model_routing; - // Get active key and decrypt - let provider_key = - match aqbot_core::repo::provider::get_active_key(&state.db, &provider.id).await { - Ok(k) => k, - Err(_) => { - return error_response( - StatusCode::BAD_GATEWAY, - &format!("No active API key for provider '{}'", provider.name), - ); - } + let targets = + match resolve_route_targets(&providers, &public_id_map, &parsed, auto_routing) { + Ok(t) => t, + Err(resp) => return resp, }; - let api_key = match decrypt_key(&provider_key.key_encrypted, &state.master_key) { - Ok(k) => k, - Err(e) => { - tracing::error!("Failed to decrypt provider key: {}", e); - return error_response(StatusCode::INTERNAL_SERVER_ERROR, "Internal key error"); - } - }; + let registry = aqbot_providers::registry::ProviderRegistry::create_default(); + let pinned = parsed.provider_hint.is_some() || targets.len() == 1 || !auto_routing; + + let mut last_error: Option = None; + for (idx, (provider, model_id)) in targets.into_iter().enumerate() { + // Get active key and decrypt + let provider_key = + match aqbot_core::repo::provider::get_active_key(&state.db, &provider.id).await { + Ok(k) => k, + Err(_) => { + let msg = format!("No active API key for provider '{}'", provider.name); + crate::auto_route::mark_failure(&provider.id, &model_id); + last_error = Some(msg.clone()); + if pinned { + return error_response(StatusCode::BAD_GATEWAY, &msg); + } + continue; + } + }; - let provider_type_str = provider_type_to_str(&provider.provider_type); + let api_key = match decrypt_key(&provider_key.key_encrypted, &state.master_key) { + Ok(k) => k, + Err(e) => { + tracing::error!("Failed to decrypt provider key: {}", e); + return error_response(StatusCode::INTERNAL_SERVER_ERROR, "Internal key error"); + } + }; - let global_settings = aqbot_core::repo::settings::get_settings(&state.db) - .await - .unwrap_or_default(); - let resolved_proxy = ProviderProxyConfig::resolve(&provider.proxy_config, &global_settings); - - let ctx = ProviderRequestContext { - api_key, - key_id: provider_key.id.clone(), - provider_id: provider.id.clone(), - base_url: Some(resolve_base_url_for_type( - &provider.api_host, - &provider.provider_type, - )), - api_path: provider.api_path.clone(), - aws_region: provider.aws_region.clone(), - proxy_config: resolved_proxy, - custom_headers: provider - .custom_headers - .as_ref() - .and_then(|s| serde_json::from_str(s).ok()), - }; + let provider_type_str = provider_type_to_str(&provider.provider_type); + let resolved_proxy = + ProviderProxyConfig::resolve(&provider.proxy_config, &global_settings); + + let ctx = ProviderRequestContext { + api_key, + key_id: provider_key.id.clone(), + provider_id: provider.id.clone(), + base_url: Some(resolve_base_url_for_type( + &provider.api_host, + &provider.provider_type, + )), + api_path: provider.api_path.clone(), + aws_region: provider.aws_region.clone(), + proxy_config: resolved_proxy, + custom_headers: provider + .custom_headers + .as_ref() + .and_then(|s| serde_json::from_str(s).ok()), + }; - let registry = aqbot_providers::registry::ProviderRegistry::create_default(); - let adapter = match registry.get(provider_type_str) { - Some(a) => a, - None => { - // Fallback to openai-compatible for custom providers - match registry.get("openai") { + let adapter = match registry.get(provider_type_str) { + Some(a) => a, + None => match registry.get("openai") { Some(a) => a, None => { - return error_response( - StatusCode::BAD_GATEWAY, - &format!("No adapter for provider type '{}'", provider_type_str), - ); + let msg = format!("No adapter for provider type '{}'", provider_type_str); + last_error = Some(msg.clone()); + if pinned { + return error_response(StatusCode::BAD_GATEWAY, &msg); + } + continue; } + }, + }; + + let mut attempt_request = request.clone(); + // Always send the real upstream model id. + attempt_request.model = model_id.clone(); + + if attempt_request.stream { + // Streaming: try until we get a stream handle; mid-stream failures + // cannot failover (handle_stream owns the response after first byte). + let response = handle_stream( + adapter, + &ctx, + attempt_request, + &state, + &gateway_key, + &provider.id, + &model_id, + start_time, + ) + .await; + // Success path returns 200 SSE; failure returns 502 JSON before stream starts. + if response.status().is_success() { + crate::auto_route::mark_success(&provider.id, &model_id); + return response; } + let msg = format!( + "Upstream '{}' failed for model '{}'", + provider.name, model_id + ); + crate::auto_route::mark_failure(&provider.id, &model_id); + last_error = Some(msg); + if pinned { + return response; + } + tracing::warn!( + attempt = idx + 1, + provider = %provider.name, + model = %model_id, + "gateway auto-route stream attempt failed; trying next" + ); + continue; } - }; - let mut request = request; - request.model = model_id.clone(); - - if request.stream { - handle_stream( + match try_non_stream( adapter, &ctx, - request, - &state, - &gateway_key, - &provider.id, - &model_id, - start_time, - ) - .await - } else { - handle_non_stream( - adapter, - &ctx, - request, + attempt_request, &state, &gateway_key, &provider.id, @@ -215,10 +295,56 @@ pub async fn chat_completions( start_time, ) .await + { + Ok(response) => { + crate::auto_route::mark_success(&provider.id, &model_id); + return response; + } + Err(err_msg) => { + let retriable = crate::auto_route::is_retriable_error_message(&err_msg); + crate::auto_route::mark_failure(&provider.id, &model_id); + last_error = Some(err_msg.clone()); + if pinned || !retriable { + let elapsed = start_time.elapsed().as_millis() as i32; + let _ = aqbot_core::repo::gateway_request_log::record_request_log( + &state.db, + &gateway_key.id, + &gateway_key.name, + "POST", + "/v1/chat/completions", + Some(&model_id), + Some(&provider.id), + 502, + elapsed, + 0, + 0, + Some(&err_msg), + ) + .await; + return error_response(StatusCode::BAD_GATEWAY, &err_msg); + } + tracing::warn!( + attempt = idx + 1, + provider = %provider.name, + model = %model_id, + error = %err_msg, + "gateway auto-route attempt failed; trying next" + ); + } + } } + + error_response( + StatusCode::BAD_GATEWAY, + last_error + .as_deref() + .unwrap_or("All upstream providers failed for this model"), + ) } -async fn handle_non_stream( +/// Attempt a non-streaming chat call. On success returns the HTTP response. +/// On failure returns the error message (caller decides failover / logging). +async fn try_non_stream( adapter: &dyn ProviderAdapter, ctx: &ProviderRequestContext, request: ChatRequest, @@ -227,10 +353,9 @@ async fn handle_non_stream( provider_id: &str, model_id: &str, start_time: Instant, -) -> axum::response::Response { +) -> Result { match adapter.chat(ctx, request).await { Ok(response) => { - // Record usage let _ = aqbot_core::repo::gateway::record_usage( &state.db, &gateway_key.id, @@ -258,28 +383,9 @@ async fn handle_non_stream( ) .await; - Json(build_non_stream_response_body(&response)).into_response() - } - Err(e) => { - let elapsed = start_time.elapsed().as_millis() as i32; - let _ = aqbot_core::repo::gateway_request_log::record_request_log( - &state.db, - &gateway_key.id, - &gateway_key.name, - "POST", - "/v1/chat/completions", - Some(model_id), - Some(provider_id), - 502, - elapsed, - 0, - 0, - Some(&e.to_string()), - ) - .await; - - error_response(StatusCode::BAD_GATEWAY, &e.to_string()) + Ok(Json(build_non_stream_response_body(&response)).into_response()) } + Err(e) => Err(e.to_string()), } } @@ -589,20 +695,38 @@ pub(crate) fn parse_model_field(model: &str, known_public_ids: &HashSet) } } -/// Resolve the `ProviderConfig` and canonical `model_id` string from a parsed -/// model field. -/// -/// - Slug hint (`/`): match enabled provider by its public ID (from the map), -/// verify model exists. -/// - No hint: scan all enabled providers for an enabled model with that ID; -/// succeed only when exactly one provider has it — otherwise error with a -/// helpful message asking the caller to use the `provider/model` form. +/// Resolve a single target (legacy helper / native pin). Prefer [`resolve_route_targets`]. +#[allow(dead_code)] pub(crate) fn resolve_provider_for_model( providers: &[ProviderConfig], public_id_map: &HashMap, parsed: &ParsedModel, ) -> Result<(ProviderConfig, String), axum::response::Response> { + let mut targets = resolve_route_targets(providers, public_id_map, parsed, false)?; + // When auto_routing is false only one target is returned; take the first. + let first = targets.remove(0); + Ok(first) +} + +/// Resolve one or more `(provider, real_model_id)` targets for a request. +/// +/// - With provider hint: always a single pinned target (model id or alias). +/// - Bare name: all matching providers (id or alias). When `auto_routing` is +/// false only the first (by sort_order) is returned; when true, the full +/// ordered pool is returned for failover. +pub(crate) fn resolve_route_targets( + providers: &[ProviderConfig], + public_id_map: &HashMap, + parsed: &ParsedModel, + auto_routing: bool, +) -> Result, axum::response::Response> { + use crate::auto_route::{ + candidates_for_attempt, collect_candidates, resolve_real_model_id, RouteCandidate, + }; + let enabled: Vec<&ProviderConfig> = providers.iter().filter(|p| p.enabled).collect(); + let provider_by_id: HashMap<&str, &ProviderConfig> = + enabled.iter().map(|p| (p.id.as_str(), *p)).collect(); match &parsed.provider_hint { Some(hint) => { @@ -617,40 +741,52 @@ pub(crate) fn resolve_provider_for_model( ) })?; - if !provider + let real_id = provider .models .iter() - .any(|m| m.enabled && m.model_id == parsed.model_id) - { + .filter(|m| m.enabled) + .find_map(|m| resolve_real_model_id(m, &parsed.model_id)) + .ok_or_else(|| { + error_response( + StatusCode::NOT_FOUND, + &format!( + "Model '{}' not found on provider '{}'", + parsed.model_id, hint + ), + ) + })?; + + Ok(vec![((*provider).clone(), real_id)]) + } + None => { + let cands = collect_candidates(providers, &parsed.model_id); + if cands.is_empty() { return Err(error_response( StatusCode::NOT_FOUND, - &format!( - "Model '{}' not found on provider '{}'", - parsed.model_id, hint - ), + &format!("Model '{}' not found", parsed.model_id), )); } - Ok(((*provider).clone(), parsed.model_id.clone())) - } - None => { - // Bare model_id: find matching enabled providers. - let matching: Vec<&&ProviderConfig> = enabled - .iter() - .filter(|p| { - p.models - .iter() - .any(|m| m.enabled && m.model_id == parsed.model_id) - }) - .collect(); + let ordered: Vec = if auto_routing && cands.len() > 1 { + candidates_for_attempt(&cands) + } else { + // First-match by sort_order (compat when auto routing is off). + cands.into_iter().take(1).collect() + }; - match matching.len() { - 0 => Err(error_response( + let mut out = Vec::with_capacity(ordered.len()); + for cand in ordered { + if let Some(provider) = provider_by_id.get(cand.provider_id.as_str()) { + out.push(((*provider).clone(), cand.real_model_id)); + } + } + if out.is_empty() { + return Err(error_response( StatusCode::NOT_FOUND, &format!("Model '{}' not found", parsed.model_id), - )), - _ => Ok(((*matching[0]).clone(), parsed.model_id.clone())), + )); } + Ok(out) } } } @@ -803,6 +939,7 @@ mod tests { param_overrides: None, image_config: None, metadata_state: None, + aliases: Vec::new(), }], ) .await @@ -971,6 +1108,7 @@ mod tests { param_overrides: None, image_config: None, metadata_state: None, + aliases: Vec::new(), }) .collect(), keys: vec![], @@ -1060,6 +1198,58 @@ mod tests { assert!(!map.contains_key(&("p2".to_string(), "gpt-4o".to_string()))); } + #[test] + fn resolve_route_targets_alias_rewrites_to_real_model_id() { + let mut providers = vec![make_provider("p1", "OpenAI", &["claude-sonnet-real"])]; + providers[0].models[0].aliases = vec!["sonnet".into()]; + let map = build_provider_public_id_map(&providers); + let parsed = parse_model_field("sonnet", &HashSet::new()); + let targets = resolve_route_targets(&providers, &map, &parsed, false).unwrap(); + assert_eq!(targets.len(), 1); + assert_eq!(targets[0].1, "claude-sonnet-real"); + } + + #[test] + fn resolve_route_targets_auto_routing_returns_pool() { + let mut providers = vec![ + make_provider("p1", "OpenAI", &["gpt-5.5"]), + make_provider("p2", "OtherAI", &["gpt-5.5"]), + ]; + providers[0].sort_order = 10; + providers[1].sort_order = 0; + let map = build_provider_public_id_map(&providers); + let parsed = parse_model_field("gpt-5.5", &HashSet::new()); + let single = resolve_route_targets(&providers, &map, &parsed, false).unwrap(); + assert_eq!(single.len(), 1); + assert_eq!(single[0].0.id, "p2"); + + let multi = resolve_route_targets(&providers, &map, &parsed, true).unwrap(); + assert_eq!(multi.len(), 2); + assert_eq!(multi[0].0.id, "p2"); + assert_eq!(multi[1].0.id, "p1"); + } + + #[test] + fn gateway_model_list_aggregates_when_auto_routing() { + let providers = vec![ + make_provider("p1", "OpenAI", &["gpt-5.5"]), + make_provider("p2", "OtherAI", &["gpt-5.5"]), + ]; + let off = build_gateway_model_list(&providers, false); + let off_ids: Vec<&str> = off.iter().filter_map(|v| v["id"].as_str()).collect(); + assert!(off_ids.iter().all(|id| id.contains('/'))); + + let on = build_gateway_model_list(&providers, true); + let on_ids: Vec<&str> = on.iter().filter_map(|v| v["id"].as_str()).collect(); + assert!(on_ids.contains(&"gpt-5.5")); + assert!(on_ids.iter().any(|id| id.contains("gpt-5.5") && id.contains('/'))); + let bare = on + .iter() + .find(|v| v["id"] == "gpt-5.5") + .expect("aggregated bare id"); + assert_eq!(bare["owned_by"], "aqbot"); + } + #[test] fn test_non_stream_payload_includes_reasoning_content() { let payload = build_non_stream_response_body(&ChatResponse { diff --git a/src-tauri/crates/gateway/src/lib.rs b/src-tauri/crates/gateway/src/lib.rs index 7ce27a42..31b6d488 100644 --- a/src-tauri/crates/gateway/src/lib.rs +++ b/src-tauri/crates/gateway/src/lib.rs @@ -1,4 +1,5 @@ pub mod auth; +pub mod auto_route; pub mod handlers; pub mod middleware; pub mod native; diff --git a/src-tauri/crates/gateway/src/native.rs b/src-tauri/crates/gateway/src/native.rs index a91b0fc4..3a7240c4 100644 --- a/src-tauri/crates/gateway/src/native.rs +++ b/src-tauri/crates/gateway/src/native.rs @@ -16,7 +16,7 @@ use tokio_stream::wrappers::ReceiverStream; use crate::{ auth::AuthenticatedKey, handlers::{ - build_provider_public_id_map, error_response, parse_model_field, resolve_provider_for_model, + build_provider_public_id_map, error_response, parse_model_field, resolve_route_targets, }, server::GatewayAppState, }; @@ -541,45 +541,40 @@ async fn resolve_native_context( )); } + let global_settings = aqbot_core::repo::settings::get_settings(&state.db) + .await + .unwrap_or_default(); + let (provider, model_id) = if let Some(model) = model { let public_id_map = build_provider_public_id_map(&candidates); let known_public_ids = public_id_map.values().cloned().collect(); let parsed = parse_model_field(model, &known_public_ids); - if parsed.provider_hint.is_some() { - let (provider, resolved_model_id) = - resolve_provider_for_model(&candidates, &public_id_map, &parsed)?; - (provider, Some(resolved_model_id)) - } else { - let matching: Vec<&ProviderConfig> = candidates - .iter() - .filter(|provider| { - provider - .models - .iter() - .any(|model| model.enabled && model.model_id == parsed.model_id) - }) - .collect(); - let fallback = matching.first().ok_or_else(|| { - error_response( - StatusCode::NOT_FOUND, - &format!("Model '{}' not found", parsed.model_id), - ) - })?; - let mut preferred_provider = None; - for provider in &matching { - if aqbot_core::repo::provider::get_active_key(&state.db, &provider.id) - .await - .is_ok() - { - preferred_provider = Some((*provider).clone()); + // Always collect the full match set so we can prefer a provider with an + // active key (compat with pre-routing behaviour). Failover across + // sources for native protocols is gated by gateway_auto_model_routing + // only after this selection path; for now we pin one upstream. + let targets = resolve_route_targets(&candidates, &public_id_map, &parsed, true)?; + + // Prefer a target with an active key when multiple are available. + // When auto routing is disabled, still prefer keys but stay on first + // healthy match by sort_order among those with keys. + let auto_routing = global_settings.gateway_auto_model_routing; + let mut selected = None; + for (provider, real_model_id) in &targets { + if aqbot_core::repo::provider::get_active_key(&state.db, &provider.id) + .await + .is_ok() + { + selected = Some((provider.clone(), real_model_id.clone())); + if !auto_routing { + // Keep first-with-key; targets are sort_order ordered. break; } + break; } - ( - preferred_provider.unwrap_or_else(|| (*fallback).clone()), - Some(parsed.model_id), - ) } + let (provider, real_model_id) = selected.unwrap_or_else(|| targets[0].clone()); + (provider, Some(real_model_id)) } else { (candidates[0].clone(), None) }; @@ -597,9 +592,6 @@ async fn resolve_native_context( error_response(StatusCode::INTERNAL_SERVER_ERROR, "Internal key error") })?; - let global_settings = aqbot_core::repo::settings::get_settings(&state.db) - .await - .unwrap_or_default(); let resolved_proxy = ProviderProxyConfig::resolve(&provider.proxy_config, &global_settings); Ok(ResolvedNativeContext { @@ -1255,6 +1247,7 @@ mod tests { param_overrides: None, image_config: None, metadata_state: None, + aliases: Vec::new(), }], ) .await @@ -1321,6 +1314,7 @@ mod tests { param_overrides: None, image_config: None, metadata_state: None, + aliases: Vec::new(), }], ) .await diff --git a/src-tauri/crates/migration/src/lib.rs b/src-tauri/crates/migration/src/lib.rs index 7eb31b63..8faef40f 100644 --- a/src-tauri/crates/migration/src/lib.rs +++ b/src-tauri/crates/migration/src/lib.rs @@ -43,6 +43,7 @@ mod m20260725_000001_add_provider_aws_region; mod m20260806_000001_add_role_capability_bindings; mod m20260807_000001_add_conversation_context_message_limit; mod m20260808_000001_compression_keep_and_source; +mod m20260809_000001_add_model_aliases_and_auto_route; pub struct Migrator; @@ -93,6 +94,7 @@ impl MigratorTrait for Migrator { Box::new(m20260806_000001_add_role_capability_bindings::Migration), Box::new(m20260807_000001_add_conversation_context_message_limit::Migration), Box::new(m20260808_000001_compression_keep_and_source::Migration), + Box::new(m20260809_000001_add_model_aliases_and_auto_route::Migration), ] } } @@ -212,7 +214,7 @@ mod tests { .expect("run sqlite migrations"); let manager = SchemaManager::new(&db); - for column in ["max_output_tokens", "metadata_state_json"] { + for column in ["max_output_tokens", "metadata_state_json", "aliases_json"] { assert!( manager .has_column("models", column) diff --git a/src-tauri/crates/migration/src/m20260809_000001_add_model_aliases_and_auto_route.rs b/src-tauri/crates/migration/src/m20260809_000001_add_model_aliases_and_auto_route.rs new file mode 100644 index 00000000..0b61cd03 --- /dev/null +++ b/src-tauri/crates/migration/src/m20260809_000001_add_model_aliases_and_auto_route.rs @@ -0,0 +1,29 @@ +use sea_orm_migration::prelude::*; + +#[derive(DeriveMigrationName)] +pub struct Migration; + +#[async_trait::async_trait] +impl MigrationTrait for Migration { + async fn up(&self, manager: &SchemaManager) -> Result<(), DbErr> { + manager + .alter_table( + Table::alter() + .table(Alias::new("models")) + .add_column(ColumnDef::new(Alias::new("aliases_json")).text().null()) + .to_owned(), + ) + .await + } + + async fn down(&self, manager: &SchemaManager) -> Result<(), DbErr> { + manager + .alter_table( + Table::alter() + .table(Alias::new("models")) + .drop_column(Alias::new("aliases_json")) + .to_owned(), + ) + .await + } +} diff --git a/src-tauri/crates/providers/src/anthropic.rs b/src-tauri/crates/providers/src/anthropic.rs index 0a93741b..c6d5d5e6 100644 --- a/src-tauri/crates/providers/src/anthropic.rs +++ b/src-tauri/crates/providers/src/anthropic.rs @@ -832,6 +832,7 @@ impl ProviderAdapter for AnthropicAdapter { param_overrides: None, image_config: None, metadata_state: None, + aliases: Vec::new(), } }) .collect()) diff --git a/src-tauri/crates/providers/src/bedrock/convert.rs b/src-tauri/crates/providers/src/bedrock/convert.rs index cef3a930..297876fa 100644 --- a/src-tauri/crates/providers/src/bedrock/convert.rs +++ b/src-tauri/crates/providers/src/bedrock/convert.rs @@ -371,6 +371,7 @@ pub(super) fn foundation_model( param_overrides: None, image_config: None, metadata_state: None, + aliases: Vec::new(), }) } diff --git a/src-tauri/crates/providers/src/cohere.rs b/src-tauri/crates/providers/src/cohere.rs index 852686f8..a017c205 100644 --- a/src-tauri/crates/providers/src/cohere.rs +++ b/src-tauri/crates/providers/src/cohere.rs @@ -57,6 +57,7 @@ pub(crate) fn cohere_models(provider_id: &str) -> Vec { param_overrides: None, image_config: None, metadata_state: None, + aliases: Vec::new(), }) .collect() } diff --git a/src-tauri/crates/providers/src/gemini.rs b/src-tauri/crates/providers/src/gemini.rs index 92aa13c4..09790bd0 100644 --- a/src-tauri/crates/providers/src/gemini.rs +++ b/src-tauri/crates/providers/src/gemini.rs @@ -787,6 +787,7 @@ impl ProviderAdapter for GeminiAdapter { param_overrides: None, image_config, metadata_state: None, + aliases: Vec::new(), } }) .collect()) diff --git a/src-tauri/crates/providers/src/jina.rs b/src-tauri/crates/providers/src/jina.rs index ad1c2fcd..2696b641 100644 --- a/src-tauri/crates/providers/src/jina.rs +++ b/src-tauri/crates/providers/src/jina.rs @@ -61,6 +61,7 @@ pub(crate) fn jina_models(provider_id: &str) -> Vec { param_overrides: None, image_config: None, metadata_state: None, + aliases: Vec::new(), }) .collect() } diff --git a/src-tauri/crates/providers/src/openai_compat.rs b/src-tauri/crates/providers/src/openai_compat.rs index eefc2c1b..0b940691 100644 --- a/src-tauri/crates/providers/src/openai_compat.rs +++ b/src-tauri/crates/providers/src/openai_compat.rs @@ -1861,6 +1861,7 @@ where param_overrides: None, image_config: None, metadata_state: None, + aliases: Vec::new(), } }) .collect() diff --git a/src-tauri/crates/providers/src/openai_responses.rs b/src-tauri/crates/providers/src/openai_responses.rs index c19eeed7..bab1de83 100644 --- a/src-tauri/crates/providers/src/openai_responses.rs +++ b/src-tauri/crates/providers/src/openai_responses.rs @@ -995,6 +995,7 @@ impl ProviderAdapter for OpenAIResponsesAdapter { param_overrides: None, image_config: None, metadata_state: None, + aliases: Vec::new(), } }) .collect()) diff --git a/src-tauri/crates/providers/src/siliconflow.rs b/src-tauri/crates/providers/src/siliconflow.rs index 3a153f4b..47d71d60 100644 --- a/src-tauri/crates/providers/src/siliconflow.rs +++ b/src-tauri/crates/providers/src/siliconflow.rs @@ -138,6 +138,7 @@ fn parse_siliconflow_image_models( param_overrides: None, image_config: None, metadata_state: None, + aliases: Vec::new(), }) .collect() } diff --git a/src-tauri/crates/providers/src/voyage.rs b/src-tauri/crates/providers/src/voyage.rs index b24d88c4..d1737d2b 100644 --- a/src-tauri/crates/providers/src/voyage.rs +++ b/src-tauri/crates/providers/src/voyage.rs @@ -57,6 +57,7 @@ pub(crate) fn voyage_models(provider_id: &str) -> Vec { param_overrides: None, image_config: None, metadata_state: None, + aliases: Vec::new(), }) .collect() } diff --git a/src-tauri/src/commands/conversations.rs b/src-tauri/src/commands/conversations.rs index 965e995a..027b53ba 100644 --- a/src-tauri/src/commands/conversations.rs +++ b/src-tauri/src/commands/conversations.rs @@ -97,6 +97,7 @@ mod function_calling_gate_tests { param_overrides: None, image_config: None, metadata_state: None, + aliases: Vec::new(), } } diff --git a/src-tauri/src/commands/drawing.rs b/src-tauri/src/commands/drawing.rs index 29c5a7ca..d0882649 100644 --- a/src-tauri/src/commands/drawing.rs +++ b/src-tauri/src/commands/drawing.rs @@ -2394,6 +2394,7 @@ mod tests { param_overrides: None, image_config: None, metadata_state: None, + aliases: Vec::new(), }], keys: Vec::new(), proxy_config: None, @@ -2445,6 +2446,7 @@ mod tests { param_overrides: None, image_config: None, metadata_state: None, + aliases: Vec::new(), }], keys: Vec::new(), proxy_config: None, @@ -2490,6 +2492,7 @@ mod tests { param_overrides: None, image_config: None, metadata_state: None, + aliases: Vec::new(), }], keys: Vec::new(), proxy_config: None, diff --git a/src-tauri/src/commands/providers.rs b/src-tauri/src/commands/providers.rs index 5a54711d..447ef5bc 100644 --- a/src-tauri/src/commands/providers.rs +++ b/src-tauri/src/commands/providers.rs @@ -835,6 +835,7 @@ mod model_metadata_tests { param_overrides: None, image_config: None, metadata_state: None, + aliases: Vec::new(), } } diff --git a/src-tauri/src/model_catalog/inference.rs b/src-tauri/src/model_catalog/inference.rs index 00090d92..23a22e0f 100644 --- a/src-tauri/src/model_catalog/inference.rs +++ b/src-tauri/src/model_catalog/inference.rs @@ -489,6 +489,7 @@ fn metadata_changes_for_new(model: &Model) -> Vec { capabilities: Vec::new(), model_type: ModelType::Chat, metadata_state: None, + aliases: Vec::new(), ..model.clone() }, model, diff --git a/src-tauri/src/model_catalog/tests/metadata.rs b/src-tauri/src/model_catalog/tests/metadata.rs index ac595f0e..85ad1d40 100644 --- a/src-tauri/src/model_catalog/tests/metadata.rs +++ b/src-tauri/src/model_catalog/tests/metadata.rs @@ -203,6 +203,7 @@ fn model(model_id: &str) -> Model { param_overrides: None, image_config: None, metadata_state: None, + aliases: Vec::new(), } } diff --git a/src/components/common/ModelParamSliders.tsx b/src/components/common/ModelParamSliders.tsx index 6bd3466a..4998a06d 100644 --- a/src/components/common/ModelParamSliders.tsx +++ b/src/components/common/ModelParamSliders.tsx @@ -56,12 +56,10 @@ function ParamRow({ step={step} value={value!} onChange={(v) => v !== null && onChange(v)} - size="small" /> )} {showSwitch && ( onChange(checked ? defaultValue : null)} /> diff --git a/src/components/gateway/GatewaySettings.tsx b/src/components/gateway/GatewaySettings.tsx index 29f17cbe..14e18038 100644 --- a/src/components/gateway/GatewaySettings.tsx +++ b/src/components/gateway/GatewaySettings.tsx @@ -211,6 +211,24 @@ export function GatewaySettings() { onChange={(checked) => handleSave({ gateway_auto_start: checked })} />
    + +
    +
    +
    + {t('gateway.autoModelRouting')} + + + +
    +
    + {t('gateway.autoModelRoutingDesc')} +
    +
    + handleSave({ gateway_auto_model_routing: checked })} + /> +
    diff --git a/src/components/settings/ImageProtocolEditor.tsx b/src/components/settings/ImageProtocolEditor.tsx index 4d6d1f1f..887a3575 100644 --- a/src/components/settings/ImageProtocolEditor.tsx +++ b/src/components/settings/ImageProtocolEditor.tsx @@ -93,7 +93,7 @@ export function ImageProtocolEditor({
    -
    + setEditAliasInput(e.target.value)} + onPressEnter={(e) => { + e.preventDefault(); + const next = editAliasInput.trim(); + if (!next || next === editingModel.model_id) { + setEditAliasInput(''); + return; + } + setEditAliases((prev) => (prev.includes(next) ? prev : [...prev, next])); + setEditAliasInput(''); + }} + onBlur={() => { + const next = editAliasInput.trim(); + if (!next || next === editingModel.model_id) { + setEditAliasInput(''); + return; + } + setEditAliases((prev) => (prev.includes(next) ? prev : [...prev, next])); + setEditAliasInput(''); + }} + style={{ flex: 1, minWidth: 120, paddingInline: 0 }} + /> +
    +
    + + + {/* Model Type */}
    @@ -2532,7 +2617,6 @@ export function ProviderDetail({ providerId }: ProviderDetailProps) {
    {t('settings.contextWindow')} { @@ -2553,8 +2637,7 @@ export function ProviderDetail({ providerId }: ProviderDetailProps) { min={1024} max={10000000} step={1024} - style={{ width: 110 }} - size="small" + style={{ width: 120 }} formatter={(value) => value ? `${Number(value).toLocaleString()}` : ''} />
    @@ -2594,7 +2677,6 @@ export function ProviderDetail({ providerId }: ProviderDetailProps) { max={10000000} placeholder={t('settings.automatic')} style={{ width: 120 }} - size="small" />
    @@ -2624,12 +2706,11 @@ export function ProviderDetail({ providerId }: ProviderDetailProps) { {/* Switches — horizontal */}
    {t('settings.useMaxCompletionTokens')} - +
    {t('settings.noSystemRole')} { setEditNoSystemRole(value); @@ -2642,7 +2723,6 @@ export function ProviderDetail({ providerId }: ProviderDetailProps) { {t('settings.omitSamplingParams')} { setEditOmitSamplingParams(value); @@ -2652,12 +2732,11 @@ export function ProviderDetail({ providerId }: ProviderDetailProps) {
    {t('settings.forceMaxTokens')} - +
    {t('settings.thinkingParamStyle')}