From b43e2fe44a9d717308f43e85a3809afddcc52350 Mon Sep 17 00:00:00 2001 From: Arepeater Date: Wed, 5 Aug 2026 17:56:27 +0800 Subject: [PATCH 01/48] =?UTF-8?q?fix(config):=20=E4=BF=AE=E5=A4=8D=20.env?= =?UTF-8?q?=20=E5=8A=A0=E8=BD=BD=E4=B8=8E=E5=AF=86=E7=A0=81=E5=88=9D?= =?UTF-8?q?=E5=A7=8B=E5=8C=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 统一开发脚本、Vite 代理与 Docker Compose 的 HOST/PORT 配置 - 首次启动按字面值读取 AUTH_PASSWORD,避免 ${VAR} 被环境变量插值 - 已有 auth.json 时继续以 Web 端密码为准,不受 .env 覆盖 --- .env.example | 2 +- Dockerfile | 3 +- backend/app/config.py | 9 ++- backend/app/services/auth.py | 10 +++- backend/tests/test_auth_env_bootstrap.py | 66 ++++++++++++++++++++++ backend/tests/test_config_env.py | 30 ++++++++++ dev.ps1 | 71 ++++++++++++++++-------- dev.sh | 49 ++++++++++++++-- docker-compose.yml | 4 +- docs/configuration.md | 9 +-- docs/deploy-password.md | 3 +- docs/deployment.md | 3 +- frontend/src/lib/api.ts | 2 +- frontend/vite.config.js | 12 ++-- frontend/vite.config.ts | 13 +++-- 15 files changed, 237 insertions(+), 49 deletions(-) create mode 100644 backend/tests/test_auth_env_bootstrap.py create mode 100644 backend/tests/test_config_env.py diff --git a/.env.example b/.env.example index 0fac35d..59db6f1 100644 --- a/.env.example +++ b/.env.example @@ -18,7 +18,7 @@ LOG_LEVEL=INFO # 首次启动时预置访问密码(可选)。公网服务器部署时填入,免去 SSH 端口转发设密码。 # 仅在尚未设置密码时生效(一次性初始化);设过后改密码请用页面 UI, 此处不再读取。 # 建议至少 6 位。.env 文件权限保持 600 且不要提交到 Git。 -AUTH_PASSWORD= +AUTH_PASSWORD='' # Optional backend dependency extras for Docker and ./dev.sh / .\dev.ps1. # Set to legacy-cpu on older CPUs without AVX2/FMA support. diff --git a/Dockerfile b/Dockerfile index cb295bf..dd0d6d6 100644 --- a/Dockerfile +++ b/Dockerfile @@ -125,7 +125,8 @@ COPY --from=stocksdk-builder /build/node_modules ./app/plugins/stocksdk/node_mod COPY tiers.yaml /app/tiers.yaml ENV STATIC_DIR=/app/static \ TIERS_YAML=/app/tiers.yaml \ - DATA_DIR=/app/data + DATA_DIR=/app/data \ + TICKFLOW_ENV_FILE=/app/.env # Frontend 静态产物 COPY --from=frontend-builder /build/dist ./static diff --git a/backend/app/config.py b/backend/app/config.py index b57f305..38be513 100644 --- a/backend/app/config.py +++ b/backend/app/config.py @@ -1,6 +1,7 @@ """全局配置 — 从环境变量 / .env 读取。""" from __future__ import annotations +import os import sys from pathlib import Path @@ -63,11 +64,17 @@ def _project_root() -> Path: _PROJECT_ROOT = _project_root() _RESOURCE_ROOT = _resource_root() +_ENV_FILE = Path( + os.environ.get( + "TICKFLOW_ENV_FILE", + str(_RESOURCE_ROOT / ".env") if not _IS_FROZEN else ".env", + ) +) class Settings(BaseSettings): model_config = SettingsConfigDict( - env_file=str(_RESOURCE_ROOT / ".env") if not _IS_FROZEN else ".env", + env_file=str(_ENV_FILE), env_file_encoding="utf-8", extra="ignore", ) diff --git a/backend/app/services/auth.py b/backend/app/services/auth.py index dbfcbad..951956a 100644 --- a/backend/app/services/auth.py +++ b/backend/app/services/auth.py @@ -131,9 +131,17 @@ def bootstrap_from_env() -> bool: Returns: True 表示本次用环境变量初始化了密码; False 表示无需初始化。 """ - from app.config import settings + from app.config import _ENV_FILE, settings pwd = (settings.auth_password or "").strip() + # Compose 会对 env_file 中未加单引号的 $VAR 做插值。Docker 部署时同时 + # 只读挂载原始 .env,首次初始化密码直接按 dotenv 语义读取,避免特殊字符被截断。 + if _ENV_FILE.is_file(): + from dotenv import dotenv_values + + raw_pwd = dotenv_values(_ENV_FILE, encoding="utf-8", interpolate=False).get("AUTH_PASSWORD") + if isinstance(raw_pwd, str) and raw_pwd.strip(): + pwd = raw_pwd.strip() if not pwd: return False if is_configured(): diff --git a/backend/tests/test_auth_env_bootstrap.py b/backend/tests/test_auth_env_bootstrap.py new file mode 100644 index 0000000..efbb41c --- /dev/null +++ b/backend/tests/test_auth_env_bootstrap.py @@ -0,0 +1,66 @@ +from __future__ import annotations + +from collections.abc import Iterator +from pathlib import Path +from types import ModuleType + +import pytest + +from app import config as app_config +from app.config import Settings + + +@pytest.fixture(autouse=True) +def isolated_auth_store( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> Iterator[tuple[ModuleType, Path, Path]]: + monkeypatch.setattr(app_config.settings, "data_dir", tmp_path) + from app.services import auth + + auth_path = tmp_path / "user_data" / "auth.json" + env_path = tmp_path / ".env" + monkeypatch.setattr(app_config, "_ENV_FILE", env_path) + monkeypatch.setattr(app_config.settings, "auth_password", "") + auth._sessions.clear() + auth._configured_cache = None + yield auth, auth_path, env_path + auth._sessions.clear() + auth._configured_cache = None + + +def test_bootstrap_recovers_compose_interpolated_password_from_raw_env( + isolated_auth_store: tuple[ModuleType, Path, Path], + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + auth, auth_path, env_path = isolated_auth_store + password = "pw${special}-secret" + env_path.write_text( + f"AUTH_PASSWORD={password}\nDATA_DIR={tmp_path}\n", + encoding="utf-8", + ) + configured = Settings(_env_file=env_path) + configured.auth_password = "pw-secret" # 模拟 Compose 将未定义的 ${special} 插值为空串 + monkeypatch.setattr(app_config, "settings", configured) + + assert auth.bootstrap_from_env() is True + assert auth_path.exists() + assert password not in auth_path.read_text(encoding="utf-8") + assert auth.verify_and_create_session(password) is not None + assert auth.verify_and_create_session("pw-secret") is None + + +def test_bootstrap_from_env_does_not_override_existing_password( + isolated_auth_store: tuple[ModuleType, Path, Path], +) -> None: + auth, auth_path, env_path = isolated_auth_store + auth.set_password("web-managed-secret") + before = auth_path.read_bytes() + env_path.write_text("AUTH_PASSWORD=replacement-secret\n", encoding="utf-8") + app_config.settings.auth_password = "replacement-secret" + + assert auth.bootstrap_from_env() is False + assert auth_path.read_bytes() == before + assert auth.verify_and_create_session("web-managed-secret") is not None + assert auth.verify_and_create_session("replacement-secret") is None diff --git a/backend/tests/test_config_env.py b/backend/tests/test_config_env.py new file mode 100644 index 0000000..1239d6e --- /dev/null +++ b/backend/tests/test_config_env.py @@ -0,0 +1,30 @@ +from __future__ import annotations + +from pathlib import Path + +import pytest + +from app.config import Settings + + +def test_settings_reads_server_and_auth_values_from_env( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + for name in ("HOST", "PORT", "LOG_LEVEL", "AUTH_PASSWORD"): + monkeypatch.delenv(name, raising=False) + env_path = tmp_path / ".env" + env_path.write_text( + "HOST=127.0.0.1\n" + "PORT=4318\n" + "LOG_LEVEL=DEBUG\n" + "AUTH_PASSWORD=config-secret\n", + encoding="utf-8", + ) + + configured = Settings(_env_file=env_path) + + assert configured.host == "127.0.0.1" + assert configured.port == 4318 + assert configured.log_level == "DEBUG" + assert configured.auth_password == "config-secret" diff --git a/dev.ps1 b/dev.ps1 index a89464f..1f9046c 100644 --- a/dev.ps1 +++ b/dev.ps1 @@ -18,8 +18,42 @@ param( $ErrorActionPreference = 'Stop' -# Port precedence: CLI arg > env var > default -if ($BackendPort -le 0) { $BackendPort = if ($env:BACKEND_PORT) { [int]$env:BACKEND_PORT } else { 3018 } } +$Root = Split-Path -Parent $MyInvocation.MyCommand.Path +$BackendDir = Join-Path $Root 'backend' +$FrontendDir = Join-Path $Root 'frontend' +$EnvFile = Join-Path $Root '.env' + +# Read only launcher-owned keys. Do not execute .env as PowerShell code. +function Read-DotEnvValue($Path, $Name) { + if (-not (Test-Path $Path)) { return $null } + $escaped = [Regex]::Escape($Name) + foreach ($line in Get-Content $Path) { + if ($line -match "^\s*$escaped\s*=\s*(.*?)\s*$") { + $value = $Matches[1].Trim() + $value = ($value -replace '\s+#.*$', '').Trim() + if ($value.Length -ge 2 -and + (($value.StartsWith('"') -and $value.EndsWith('"')) -or + ($value.StartsWith("'") -and $value.EndsWith("'")))) { + return $value.Substring(1, $value.Length - 2) + } + return $value + } + } + return $null +} + +$DotEnvHost = Read-DotEnvValue $EnvFile 'HOST' +$DotEnvPort = Read-DotEnvValue $EnvFile 'PORT' +$BindAddress = if ($env:HOST) { $env:HOST } elseif ($DotEnvHost) { $DotEnvHost } else { '0.0.0.0' } +$DisplayHost = if ($BindAddress -in @('0.0.0.0', '::')) { 'localhost' } else { $BindAddress } + +# Port precedence: CLI arg > BACKEND_PORT env > PORT env > .env PORT > default +if ($BackendPort -le 0) { + if ($env:BACKEND_PORT) { $BackendPort = [int]$env:BACKEND_PORT } + elseif ($env:PORT) { $BackendPort = [int]$env:PORT } + elseif ($DotEnvPort) { $BackendPort = [int]$DotEnvPort } + else { $BackendPort = 3018 } +} if ($FrontendPort -le 0) { $FrontendPort = if ($env:FRONTEND_PORT) { [int]$env:FRONTEND_PORT } else { 3011 } } # Force UTF-8 console output so child process logs aren't garbled @@ -28,10 +62,6 @@ try { $OutputEncoding = New-Object System.Text.UTF8Encoding $false } catch {} -$Root = Split-Path -Parent $MyInvocation.MyCommand.Path -$BackendDir = Join-Path $Root 'backend' -$FrontendDir = Join-Path $Root 'frontend' - function Log-Info($m) { Write-Host "[dev] $m" -ForegroundColor DarkGray } function Log-Ok ($m) { Write-Host "[dev] $m" -ForegroundColor Green } function Log-Warn($m) { Write-Host "[dev] $m" -ForegroundColor Yellow } @@ -113,15 +143,7 @@ Free-Port 'frontend' $FrontendPort # select Polars' rtcompat runtime before the backend starts. $BackendExtras = $env:BACKEND_EXTRAS if (-not (Test-Path Env:BACKEND_EXTRAS)) { - $envFile = Join-Path $Root '.env' - if (Test-Path $envFile) { - foreach ($line in Get-Content $envFile) { - if ($line -match '^\s*BACKEND_EXTRAS\s*=\s*(.*?)\s*$') { - $BackendExtras = $Matches[1] - break - } - } - } + $BackendExtras = Read-DotEnvValue $EnvFile 'BACKEND_EXTRAS' } $BackendExtraArgs = @() @@ -156,8 +178,8 @@ Write-Host '' Write-Host '+----------------------------------------------+' -ForegroundColor Blue Write-Host '| tickflow-stock-panel |' -ForegroundColor Blue Write-Host '| |' -ForegroundColor Blue -Write-Host "| backend http://localhost:$BackendPort" -ForegroundColor Blue -Write-Host "| frontend http://localhost:$FrontendPort" -ForegroundColor Blue +Write-Host "| backend http://${DisplayHost}:$BackendPort" -ForegroundColor Blue +Write-Host "| frontend http://${DisplayHost}:$FrontendPort" -ForegroundColor Blue Write-Host '| |' -ForegroundColor Blue Write-Host '| Ctrl-C closes both |' -ForegroundColor Blue Write-Host '+----------------------------------------------+' -ForegroundColor Blue @@ -170,7 +192,7 @@ $backendPidFile = [System.IO.Path]::GetTempFileName() $frontendPidFile = [System.IO.Path]::GetTempFileName() $backendJob = Start-Job -Name 'backend' -ScriptBlock { - param($pidFile, $dir, $port) + param($pidFile, $dir, $envFile, $bindAddress, $port) # Start-Job 开的是全新 powershell.exe 子进程, 不继承主进程的 UTF-8 设置, # 默认用系统 ANSI (中文 Windows = GBK/cp936) 解码后端 UTF-8 输出 → 中文乱码。 # 这里强制子进程用 UTF-8, 与 app/__init__.py 的 stdout/stderr 编码对齐。 @@ -179,18 +201,21 @@ $backendJob = Start-Job -Name 'backend' -ScriptBlock { $PID | Out-File -FilePath $pidFile -Encoding ascii -Force $env:PYTHONUNBUFFERED = '1' Set-Location $dir - & .\.venv\Scripts\python.exe -m uvicorn app.main:app --reload --host 0.0.0.0 --port $port 2>&1 -} -ArgumentList $backendPidFile, $BackendDir, $BackendPort + $envArgs = if (Test-Path $envFile) { @('--env-file', $envFile) } else { @() } + & .\.venv\Scripts\python.exe -m uvicorn app.main:app @envArgs --reload --host $bindAddress --port $port 2>&1 +} -ArgumentList $backendPidFile, $BackendDir, $EnvFile, $BindAddress, $BackendPort $frontendJob = Start-Job -Name 'frontend' -ScriptBlock { - param($pidFile, $dir, $port) + param($pidFile, $dir, $bindAddress, $backendPort, $port) # 同上: job 子进程默认 GBK, pnpm/前端工具链也是 UTF-8 输出, 需对齐。 [Console]::OutputEncoding = New-Object System.Text.UTF8Encoding $false $OutputEncoding = New-Object System.Text.UTF8Encoding $false $PID | Out-File -FilePath $pidFile -Encoding ascii -Force Set-Location $dir - & pnpm dev --host 0.0.0.0 --port $port 2>&1 -} -ArgumentList $frontendPidFile, $FrontendDir, $FrontendPort + $env:BACKEND_HOST = $bindAddress + $env:BACKEND_PORT = [string]$backendPort + & pnpm dev --host $bindAddress --port $port 2>&1 +} -ArgumentList $frontendPidFile, $FrontendDir, $BindAddress, $BackendPort, $FrontendPort # Wait up to 5 seconds for the PID files to materialise function Read-JobPid($file) { diff --git a/dev.sh b/dev.sh index b8deca7..09f8a67 100755 --- a/dev.sh +++ b/dev.sh @@ -13,13 +13,48 @@ set -euo pipefail ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" BACKEND_DIR="$ROOT/backend" FRONTEND_DIR="$ROOT/frontend" -BACKEND_PORT="${BACKEND_PORT:-3018}" + +# Read only the launcher-owned keys from .env. Do not source the whole file: +# .env is data, not a shell script, and may contain values that are unsafe or +# invalid as Bash syntax. Exported environment variables keep highest priority. +read_dotenv_value() { + local key="$1" + if [[ ! -f "$ROOT/.env" ]]; then + return 0 + fi + awk -v wanted="$key" ' + $0 ~ "^[[:space:]]*" wanted "[[:space:]]*=" { + sub(/^[^=]*=/, "") + sub(/[[:space:]]+#.*$/, "") + gsub(/^[[:space:]]+|[[:space:]]+$/, "") + if (($0 ~ /^".*"$/) || ($0 ~ /^\047.*\047$/)) { + $0 = substr($0, 2, length($0) - 2) + } + print + exit + } + ' "$ROOT/.env" +} + +ENV_HOST="$(read_dotenv_value HOST)" +ENV_PORT="$(read_dotenv_value PORT)" +BACKEND_HOST="${HOST:-${ENV_HOST:-0.0.0.0}}" +# Keep BACKEND_PORT as a backwards-compatible explicit override. +BACKEND_PORT="${BACKEND_PORT:-${PORT:-${ENV_PORT:-3018}}}" FRONTEND_PORT="${FRONTEND_PORT:-3011}" +UVICORN_ENV_ARGS=() +if [[ -f "$ROOT/.env" ]]; then + UVICORN_ENV_ARGS=(--env-file "$ROOT/.env") +fi +DISPLAY_HOST="$BACKEND_HOST" +if [[ "$DISPLAY_HOST" == "0.0.0.0" || "$DISPLAY_HOST" == "::" ]]; then + DISPLAY_HOST="localhost" +fi # Match Docker's BACKEND_EXTRAS behavior so old CPUs can select Polars' # rtcompat runtime before the backend starts. An exported value wins over .env. if [[ -z "${BACKEND_EXTRAS+x}" && -f "$ROOT/.env" ]]; then - BACKEND_EXTRAS="$(awk '/^[[:space:]]*BACKEND_EXTRAS[[:space:]]*=/ {sub(/^[^=]*=/, ""); gsub(/^[[:space:]]+|[[:space:]]+$/, ""); print; exit}' "$ROOT/.env")" + BACKEND_EXTRAS="$(read_dotenv_value BACKEND_EXTRAS)" fi BACKEND_EXTRAS="${BACKEND_EXTRAS:-}" BACKEND_EXTRA_ARGS=() @@ -129,8 +164,8 @@ echo echo -e "${BLUE}╭──────────────────────────────────────────────╮${NC}" echo -e "${BLUE}│${NC} ${GREEN}tickflow-stock-panel${NC} ${BLUE}│${NC}" echo -e "${BLUE}│${NC} ${BLUE}│${NC}" -echo -e "${BLUE}│${NC} backend ${YELLOW}http://localhost:$BACKEND_PORT${NC} ${BLUE}│${NC}" -echo -e "${BLUE}│${NC} frontend ${YELLOW}http://localhost:$FRONTEND_PORT${NC} ${BLUE}│${NC}" +echo -e "${BLUE}│${NC} backend ${YELLOW}http://$DISPLAY_HOST:$BACKEND_PORT${NC} ${BLUE}│${NC}" +echo -e "${BLUE}│${NC} frontend ${YELLOW}http://$DISPLAY_HOST:$FRONTEND_PORT${NC} ${BLUE}│${NC}" echo -e "${BLUE}│${NC} ${BLUE}│${NC}" echo -e "${BLUE}│${NC} Ctrl-C 同时关闭两端 ${BLUE}│${NC}" echo -e "${BLUE}╰──────────────────────────────────────────────╯${NC}" @@ -138,14 +173,16 @@ echo ( cd "$BACKEND_DIR" - uv run uvicorn app.main:app --reload --host 0.0.0.0 --port "$BACKEND_PORT" 2>&1 \ + uv run uvicorn app.main:app "${UVICORN_ENV_ARGS[@]}" --reload \ + --host "$BACKEND_HOST" --port "$BACKEND_PORT" 2>&1 \ | prefix_awk "$(printf "${BLUE}[backend ]${NC} ")" ) & PIDS+=("$!") ( cd "$FRONTEND_DIR" - pnpm dev --host 0.0.0.0 --port "$FRONTEND_PORT" 2>&1 \ + BACKEND_HOST="$BACKEND_HOST" BACKEND_PORT="$BACKEND_PORT" \ + pnpm dev --host "$BACKEND_HOST" --port "$FRONTEND_PORT" 2>&1 \ | prefix_awk "$(printf "${GREEN}[frontend]${NC} ")" ) & PIDS+=("$!") diff --git a/docker-compose.yml b/docker-compose.yml index f705549..c4ed028 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -10,7 +10,7 @@ services: CODEX_CLI_VERSION: ${CODEX_CLI_VERSION:-0.144.3} container_name: TickFlow_Stock_Panel ports: - - "${PORT:-3018}:3018" + - "${HOST:-0.0.0.0}:${PORT:-3018}:3018" extra_hosts: - "host.docker.internal:host-gateway" env_file: @@ -26,6 +26,8 @@ services: volumes: - ./data:/app/data - ./tiers.yaml:/app/tiers.yaml:ro + # 保留原始 dotenv 值供首次密码初始化读取,避免 Compose 展开密码中的 $VAR。 + - ./.env:/app/.env:ro # 复用主机 Codex 登录态;后端只读后复制到单次请求的临时 CODEX_HOME。 # Windows PowerShell/CMD 下 HOME 常未设置, 可通过 .env 里 CODEX_HOME_HOST 覆盖。 - ${CODEX_HOME_HOST:-${HOME}/.codex}:/root/.codex:ro diff --git a/docs/configuration.md b/docs/configuration.md index ca774cf..21c38ad 100644 --- a/docs/configuration.md +++ b/docs/configuration.md @@ -58,13 +58,13 @@ AI_DAILY_TOKEN_BUDGET=500000 # 每日 token 预算上限 ## 服务 ```ini -HOST=0.0.0.0 # 监听地址 -PORT=3018 # 服务端口 +HOST=0.0.0.0 # 开发服务监听地址 / Docker 主机绑定地址 +PORT=3018 # 开发后端端口 / Docker 主机映射端口 LOG_LEVEL=INFO # DEBUG | INFO | WARNING | ERROR ``` - `HOST`:`0.0.0.0` 监听所有网卡(容器/公网部署需要);仅本机用可设 `127.0.0.1` -- `PORT`:默认 `3018`,改端口后 Docker 映射、SSH 转发命令里的端口也要同步改 +- `PORT`:默认 `3018`;开发模式兼容显式的 `BACKEND_PORT` 覆盖,改端口后 SSH 转发命令也要同步改 - `LOG_LEVEL`:排查问题时改 `DEBUG` --- @@ -84,10 +84,11 @@ DATA_DIR=./data # Parquet / DuckDB 数据存储目录 ## 访问密码(公网部署) ```ini -AUTH_PASSWORD=你的密码 # 至少 6 位;仅首次生效,已设过则不覆盖 +AUTH_PASSWORD='你的密码' # 至少 6 位;仅首次生效,已设过则不覆盖 ``` 面板首次设置访问密码时,出于安全考虑**仅允许本机或内网访问**(防公网陌生人抢先设置锁死面板)。公网服务器部署可通过此环境变量预置首个密码。 +密码建议使用单引号包裹,避免 Docker Compose 插值 `$VAR`;Docker 启动时也会只读挂载原始 `.env`,兼容已有的未加引号配置。 详细步骤、SSH 转发方案、重置密码方法见 [deployment.md → 访问密码设置](./deployment.md#访问密码设置公网部署必读)。 diff --git a/docs/deploy-password.md b/docs/deploy-password.md index c828b31..9a1c053 100644 --- a/docs/deploy-password.md +++ b/docs/deploy-password.md @@ -16,7 +16,7 @@ ```bash # 编辑服务器上的 .env (通常在项目根目录或 backend/ 下) -AUTH_PASSWORD=你的密码 +AUTH_PASSWORD='你的密码' ``` 然后重启服务。启动时会自动: @@ -30,6 +30,7 @@ AUTH_PASSWORD=你的密码 - **密码至少 6 位**,否则会被跳过并记一条 warning 日志 - **仅在未设过密码时生效**。已设过密码后,改这里不会覆盖(避免重启时重置你在 UI 改的密码) +- 密码建议使用单引号包裹,避免 Docker Compose 插值 `$VAR`;启动时也会从只读挂载的原始 `.env` 初始化,兼容已有的未加引号配置 - `.env` 文件权限保持 `600`,**不要提交到 Git** - 明文密码只存在于 `.env` / 环境变量中,落盘的是哈希,安全性等同 `auth.json` diff --git a/docs/deployment.md b/docs/deployment.md index 8e006a5..af16e40 100644 --- a/docs/deployment.md +++ b/docs/deployment.md @@ -123,7 +123,7 @@ git pull 在 `.env` 文件(或 Docker / 系统环境变量)里设置 `AUTH_PASSWORD`: ```bash -AUTH_PASSWORD=你的密码 +AUTH_PASSWORD='你的密码' ``` 然后重启服务。启动时会自动: @@ -138,6 +138,7 @@ AUTH_PASSWORD=你的密码 - **密码至少 6 位**,否则会被跳过并记一条 warning 日志 - **仅在未设过密码时生效**。已设过密码后,改这里不会覆盖(避免重启时重置你在 UI 改的密码) +- 密码建议使用单引号包裹,避免 Docker Compose 插值 `$VAR`;启动时也会从只读挂载的原始 `.env` 初始化,兼容已有的未加引号配置 - `.env` 文件权限保持 `600`,**不要提交到 Git** - 明文密码只存在于 `.env` / 环境变量中,落盘的是哈希,安全性等同 `auth.json` diff --git a/frontend/src/lib/api.ts b/frontend/src/lib/api.ts index 7b6f8d9..0d03ab9 100644 --- a/frontend/src/lib/api.ts +++ b/frontend/src/lib/api.ts @@ -1,6 +1,6 @@ // 后端 API 客户端 — 全项目统一入口 // -// Dev:Vite 代理 /api 到 :3018 +// Dev: Vite 按启动脚本解析出的 BACKEND_HOST/BACKEND_PORT 代理 /api // Prod:同源(FastAPI 托管前端 dist) import { toast } from '@/components/Toast' diff --git a/frontend/vite.config.js b/frontend/vite.config.js index 50f463f..239f479 100644 --- a/frontend/vite.config.js +++ b/frontend/vite.config.js @@ -1,6 +1,10 @@ import { defineConfig } from 'vite'; import react from '@vitejs/plugin-react'; import path from 'node:path'; +const backendHost = process.env.BACKEND_HOST || '127.0.0.1'; +const proxyHost = ['0.0.0.0', '::'].includes(backendHost) ? '127.0.0.1' : backendHost; +const backendPort = process.env.BACKEND_PORT || '3018'; +const backendTarget = `http://${proxyHost}:${backendPort}`; export default defineConfig({ plugins: [react()], resolve: { @@ -9,12 +13,12 @@ export default defineConfig({ }, }, server: { - host: '0.0.0.0', // 允许局域网访问 + host: '0.0.0.0', // dev.sh / dev.ps1 会用 CLI --host 覆盖 port: 3011, proxy: { - // dev 时 /api 转发到 FastAPI + // dev 时 /api 转发到与启动脚本相同的 FastAPI 地址 '/api': { - target: 'http://localhost:3018', + target: backendTarget, // SSE 端点需要禁用缓冲 configure: (proxy) => { proxy.on('proxyReq', (_proxyReq, req) => { @@ -26,7 +30,7 @@ export default defineConfig({ }); }, }, - '/health': 'http://localhost:3018', + '/health': backendTarget, }, }, build: { diff --git a/frontend/vite.config.ts b/frontend/vite.config.ts index 554d93d..51d5a55 100644 --- a/frontend/vite.config.ts +++ b/frontend/vite.config.ts @@ -2,6 +2,11 @@ import { defineConfig } from 'vite' import react from '@vitejs/plugin-react' import path from 'node:path' +const backendHost = process.env.BACKEND_HOST || '127.0.0.1' +const proxyHost = ['0.0.0.0', '::'].includes(backendHost) ? '127.0.0.1' : backendHost +const backendPort = process.env.BACKEND_PORT || '3018' +const backendTarget = `http://${proxyHost}:${backendPort}` + export default defineConfig({ plugins: [react()], resolve: { @@ -10,12 +15,12 @@ export default defineConfig({ }, }, server: { - host: '0.0.0.0', // 允许局域网访问 + host: '0.0.0.0', // dev.sh / dev.ps1 会用 CLI --host 覆盖 port: 3011, proxy: { - // dev 时 /api 转发到 FastAPI + // dev 时 /api 转发到与启动脚本相同的 FastAPI 地址 '/api': { - target: 'http://localhost:3018', + target: backendTarget, // SSE 端点需要禁用缓冲 configure: (proxy) => { proxy.on('proxyReq', (_proxyReq, req) => { @@ -27,7 +32,7 @@ export default defineConfig({ }) }, }, - '/health': 'http://localhost:3018', + '/health': backendTarget, }, }, build: { From 7dad95c4f1074b77d9fbcef475227800d3c75943 Mon Sep 17 00:00:00 2001 From: Arepeater Date: Wed, 5 Aug 2026 17:57:05 +0800 Subject: [PATCH 02/48] =?UTF-8?q?feat(settings):=20=E6=94=AF=E6=8C=81?= =?UTF-8?q?=E9=85=8D=E7=BD=AE=E6=95=B0=E6=8D=AE=E4=BB=BB=E5=8A=A1=E8=B6=85?= =?UTF-8?q?=E6=97=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 在数据源设置页配置普通任务与分钟 K 长任务超时 - 支持秒、分钟、小时输入,持久化时统一换算为秒 - 保留 1200/1800 秒默认值,配置仅影响新创建的任务 --- backend/app/api/kline.py | 4 +- backend/app/api/settings.py | 15 ++ backend/app/services/pipeline_jobs.py | 21 ++- backend/app/services/preferences.py | 26 ++++ .../tests/test_pipeline_and_monitor_fixes.py | 12 +- frontend/src/lib/api.ts | 13 ++ frontend/src/pages/settings/DataSources.tsx | 144 +++++++++++++++++- 7 files changed, 223 insertions(+), 12 deletions(-) diff --git a/backend/app/api/kline.py b/backend/app/api/kline.py index 5dcde88..72b775c 100644 --- a/backend/app/api/kline.py +++ b/backend/app/api/kline.py @@ -813,7 +813,7 @@ async def sync_minute(request: Request): """ import asyncio - from app.services.pipeline_jobs import job_store, release_run_slot, try_acquire_run_slot, LONG_JOB_TIMEOUT_S + from app.services.pipeline_jobs import job_store, release_run_slot, try_acquire_run_slot from app.api.data import invalidate_storage_cache from app.services.preferences import get_minute_sync_days from app.tickflow.capabilities import Cap @@ -836,7 +836,7 @@ async def sync_minute(request: Request): extend_flag = body.get("extend") # 分钟K全市场同步是长任务(数据量是日K的 ~240 倍),用更宽松的卡死阈值 - job_id, is_new = job_store.create(timeout_s=LONG_JOB_TIMEOUT_S) + job_id, is_new = job_store.create(long_running=True) if not is_new: return {"status": "reused", "job_id": job_id} diff --git a/backend/app/api/settings.py b/backend/app/api/settings.py index 6528466..55f8bdf 100644 --- a/backend/app/api/settings.py +++ b/backend/app/api/settings.py @@ -350,6 +350,11 @@ class DataProvidersIn(BaseModel): financial_data_provider: str | None = None +class DataSourceJobTimeoutPrefs(BaseModel): + data_source_job_timeout_s: int = Field(ge=60) + data_source_long_job_timeout_s: int = Field(ge=60) + + class DatasetFieldMapItem(BaseModel): source: str target: str @@ -413,6 +418,8 @@ def get_preferences() -> dict: "minute_data_provider": preferences.get_minute_data_provider(), "realtime_data_provider": preferences.get_realtime_data_provider(), "financial_data_provider": preferences.get_financial_provider(), + "data_source_job_timeout_s": preferences.get_data_source_job_timeout_s(), + "data_source_long_job_timeout_s": preferences.get_data_source_long_job_timeout_s(), "realtime_watchlist_symbols": preferences.get_realtime_watchlist_symbols(), **preferences.get_realtime_quote_scope(), "pipeline_pull_a_share": preferences.get_pipeline_pull_a_share(), @@ -616,6 +623,14 @@ def update_data_providers(req: DataProvidersIn) -> dict: } +@router.put("/preferences/data-source-job-timeouts") +def update_data_source_job_timeouts(req: DataSourceJobTimeoutPrefs) -> dict: + """保存普通与长数据后台任务的卡死判定时间。""" + from app.services import preferences + preferences.save(req.model_dump()) + return req.model_dump() + + @router.get("/preferences/watchlist-columns") def get_watchlist_columns() -> dict: """返回自选列表列配置。""" diff --git a/backend/app/services/pipeline_jobs.py b/backend/app/services/pipeline_jobs.py index 53f3397..c4b99eb 100644 --- a/backend/app/services/pipeline_jobs.py +++ b/backend/app/services/pipeline_jobs.py @@ -27,7 +27,7 @@ JobStatus = Literal["pending", "running", "succeeded", "failed"] # 由 reap_stale() 在 /run 和 /jobs/{id} 轮询端点检查 — 保证卡死后能自愈, # 无需用户再次点击「同步」。 # -# 超时阈值按任务类型区分: +# 默认超时阈值按任务类型区分,可在 Web 数据源设置中调整: # - 普通任务(日K管道/扩展/修正/重算): 1200s (20 分钟) # - 长任务(分钟K全市场同步,数据量是日K的 ~240 倍): 1800s (30 分钟) # 分钟K即使流式落盘后仍可能跑十几到数十分钟(限速 sleep 是主因), @@ -105,7 +105,12 @@ class JobStore: # ===== lifecycle ===== - def create(self, timeout_s: int = DEFAULT_JOB_TIMEOUT_S) -> tuple[str, bool]: + def create( + self, + timeout_s: int | None = None, + *, + long_running: bool = False, + ) -> tuple[str, bool]: """单飞创建任务。返回 (job_id, is_new)。 去重条件为 **pending ∨ running**(而非仅 running):`/run` 先 create() 再在 @@ -115,9 +120,17 @@ class JobStore: is_new=False 表示复用了已有活跃任务,调用方**不得**再调度新的后台任务。 - timeout_s: reap_stale 判定卡死的阈值。普通任务默认 1200s; - 分钟K全市场同步等长任务传 LONG_JOB_TIMEOUT_S (1800s)。 + timeout_s: reap_stale 判定卡死的阈值。None 时读取用户配置。 + long_running: timeout_s 为 None 时,是否读取长任务配置;普通任务默认 + 1200s,分钟K全市场同步等长任务默认 1800s。 """ + if timeout_s is None: + from app.services import preferences + if long_running: + timeout_s = preferences.get_data_source_long_job_timeout_s() + else: + timeout_s = preferences.get_data_source_job_timeout_s() + with self._lock: if self._active_id: active = self._active_jobs.get(self._active_id) diff --git a/backend/app/services/preferences.py b/backend/app/services/preferences.py index 5671bf8..42d47af 100644 --- a/backend/app/services/preferences.py +++ b/backend/app/services/preferences.py @@ -186,6 +186,32 @@ def get_minute_sync_segment_days() -> int: # ===== 数据源选择 (默认 TickFlow;第一阶段仅日K切换入口) ===== _ALLOWED_DATA_PROVIDERS = {"tickflow"} +DATA_SOURCE_JOB_TIMEOUT_MIN_S = 60 + + +def get_data_source_job_timeout_s() -> int: + """返回普通数据后台任务的卡死判定时间(秒)。""" + from app.services.pipeline_jobs import DEFAULT_JOB_TIMEOUT_S + raw = load().get("data_source_job_timeout_s", DEFAULT_JOB_TIMEOUT_S) + try: + timeout_s = int(raw) + except (TypeError, ValueError): + timeout_s = DEFAULT_JOB_TIMEOUT_S + return max(DATA_SOURCE_JOB_TIMEOUT_MIN_S, timeout_s) + + +def get_data_source_long_job_timeout_s() -> int: + """返回分钟 K 全市场等长任务的卡死判定时间(秒)。""" + from app.services.pipeline_jobs import LONG_JOB_TIMEOUT_S + raw = load().get( + "data_source_long_job_timeout_s", + LONG_JOB_TIMEOUT_S, + ) + try: + timeout_s = int(raw) + except (TypeError, ValueError): + timeout_s = LONG_JOB_TIMEOUT_S + return max(DATA_SOURCE_JOB_TIMEOUT_MIN_S, timeout_s) def _allowed_data_providers() -> set[str]: diff --git a/backend/tests/test_pipeline_and_monitor_fixes.py b/backend/tests/test_pipeline_and_monitor_fixes.py index a3dd75e..3b3615c 100644 --- a/backend/tests/test_pipeline_and_monitor_fixes.py +++ b/backend/tests/test_pipeline_and_monitor_fixes.py @@ -8,7 +8,7 @@ import polars as pl import pytest from app.jobs import daily_pipeline -from app.services import pipeline_jobs, quote_service +from app.services import pipeline_jobs, preferences, quote_service from app.services.pipeline_jobs import JobStore from app.services.quote_service import QuoteService from app.strategy import monitor_rules @@ -16,12 +16,14 @@ from app.strategy.monitor import MonitorRuleEngine # ── JobStore 单飞 ──────────────────────────────────────────────────────── -def test_create_singleflight_dedupes_pending_window(tmp_path): +def test_create_singleflight_dedupes_pending_window(monkeypatch, tmp_path): """两次快速 create() 在 pending 窗口内应复用同一 job(is_new=False)。""" + monkeypatch.setattr(preferences, "load", lambda: {"data_source_job_timeout_s": 3600}) store = JobStore(store_dir=tmp_path / "jobs") jid1, new1 = store.create() assert new1 is True + assert store.get(jid1)["timeout_s"] == 3600 # 尚未 start(), job 仍是 pending —— 旧实现会在此另起新 job(并发双跑根因) jid2, new2 = store.create() @@ -35,10 +37,12 @@ def test_create_singleflight_dedupes_pending_window(tmp_path): assert new3 is False -def test_create_new_after_terminal(tmp_path): +def test_create_new_after_terminal(monkeypatch, tmp_path): """job 终态(succeed/fail)后, create() 应给出新 job。""" + monkeypatch.setattr(preferences, "load", lambda: {"data_source_long_job_timeout_s": 5400}) store = JobStore(store_dir=tmp_path / "jobs") - jid1, _ = store.create() + jid1, _ = store.create(long_running=True) + assert store.get(jid1)["timeout_s"] == 5400 store.start(jid1) store.succeed(jid1, {"ok": True}) diff --git a/frontend/src/lib/api.ts b/frontend/src/lib/api.ts index 0d03ab9..c4f2e27 100644 --- a/frontend/src/lib/api.ts +++ b/frontend/src/lib/api.ts @@ -969,6 +969,8 @@ export interface Preferences { minute_data_provider?: string realtime_data_provider?: string financial_data_provider?: string + data_source_job_timeout_s: number + data_source_long_job_timeout_s: number realtime_watchlist_symbols?: string[] realtime_pull_stock?: boolean realtime_pull_etf?: boolean @@ -1126,6 +1128,17 @@ export const api = { '/api/settings/preferences/data-providers', { method: 'PUT', body: JSON.stringify(cfg) }, ), + updateDataSourceJobTimeouts: (dataSourceJobTimeoutS: number, dataSourceLongJobTimeoutS: number) => + request>( + '/api/settings/preferences/data-source-job-timeouts', + { + method: 'PUT', + body: JSON.stringify({ + data_source_job_timeout_s: dataSourceJobTimeoutS, + data_source_long_job_timeout_s: dataSourceLongJobTimeoutS, + }), + }, + ), updateMinuteSync: (enabled: boolean, days: number, segmentDays?: number) => request('/api/settings/preferences/minute-sync', { method: 'PUT', diff --git a/frontend/src/pages/settings/DataSources.tsx b/frontend/src/pages/settings/DataSources.tsx index dc57b08..9744fa3 100644 --- a/frontend/src/pages/settings/DataSources.tsx +++ b/frontend/src/pages/settings/DataSources.tsx @@ -1,8 +1,8 @@ import { useState } from 'react' import { useMutation, useQuery, useQueryClient } from '@tanstack/react-query' import { motion, AnimatePresence } from 'framer-motion' -import { Check, Database, Plus, RefreshCw, Zap, FileWarning } from 'lucide-react' -import { api, type DataSourceItem, type PluginDataSourceItem } from '@/lib/api' +import { Check, Clock3, Database, Plus, RefreshCw, Zap, FileWarning } from 'lucide-react' +import { api, type DataSourceItem, type PluginDataSourceItem, type Preferences } from '@/lib/api' import { QK } from '@/lib/queryKeys' import { usePreferences } from '@/lib/useSharedQueries' import { toast } from '@/components/Toast' @@ -15,12 +15,65 @@ const DATASET_LABEL: Record = { minute: '分钟', } +type TimeoutUnit = 'second' | 'minute' | 'hour' + +const TIMEOUT_UNIT_SECONDS: Record = { + second: 1, + minute: 60, + hour: 3600, +} + +function preferredTimeoutUnit(seconds: number): TimeoutUnit { + if (seconds >= 3600 && seconds % 1800 === 0) return 'hour' + if (seconds % 60 === 0) return 'minute' + return 'second' +} + +function formatTimeoutValue(seconds: number, unit: TimeoutUnit): string { + if (!Number.isFinite(seconds)) return '' + const value = seconds / TIMEOUT_UNIT_SECONDS[unit] + return String(Number(value.toFixed(4))) +} + export function SettingsDataSourcesPanel() { const qc = useQueryClient() const prefs = usePreferences() const sources = useQuery({ queryKey: QK.dataSources, queryFn: api.dataSources }) const [selected, setSelected] = useState('tickflow') // 当前在右侧编辑的源 name const [confirmDelete, setConfirmDelete] = useState(null) + const [timeoutDraft, setTimeoutDraft] = useState<{ regular: string; long: string } | null>(null) + const [regularUnitOverride, setRegularUnitOverride] = useState(null) + const [longUnitOverride, setLongUnitOverride] = useState(null) + + const currentRegularTimeout = prefs.data?.data_source_job_timeout_s ?? 1200 + const currentLongTimeout = prefs.data?.data_source_long_job_timeout_s ?? 1800 + const regularTimeoutUnit = regularUnitOverride ?? preferredTimeoutUnit(currentRegularTimeout) + const longTimeoutUnit = longUnitOverride ?? preferredTimeoutUnit(currentLongTimeout) + const regularTimeoutInput = timeoutDraft?.regular + ?? formatTimeoutValue(currentRegularTimeout, regularTimeoutUnit) + const longTimeoutInput = timeoutDraft?.long + ?? formatTimeoutValue(currentLongTimeout, longTimeoutUnit) + const regularInputNumber = Number(regularTimeoutInput) + const longInputNumber = Number(longTimeoutInput) + const regularTimeout = Math.round(regularInputNumber * TIMEOUT_UNIT_SECONDS[regularTimeoutUnit]) + const longTimeout = Math.round(longInputNumber * TIMEOUT_UNIT_SECONDS[longTimeoutUnit]) + const timeoutValuesValid = Number.isFinite(regularInputNumber) && regularInputNumber > 0 + && Number.isFinite(longInputNumber) && longInputNumber > 0 + && regularTimeout >= 60 && longTimeout >= 60 + const timeoutValuesChanged = regularTimeout !== currentRegularTimeout + || longTimeout !== currentLongTimeout + + const saveJobTimeouts = useMutation({ + mutationFn: () => api.updateDataSourceJobTimeouts(regularTimeout, longTimeout), + onSuccess: (saved) => { + qc.setQueryData(QK.preferences, current => ( + current ? { ...current, ...saved } : current + )) + setTimeoutDraft(null) + toast('任务超时配置已保存', 'success') + }, + onError: (e: Error) => toast(`保存失败: ${e.message}`, 'error'), + }) const reload = useMutation({ mutationFn: api.reloadDataSources, @@ -312,6 +365,93 @@ export function SettingsDataSourcesPanel() { +
+
+
+ +
+

数据任务超时

+

+ 后台任务运行超过对应时间后将判定为疑似卡死。保存时自动换算为秒,修改后对新建任务生效。 +

+
+
+ +
+ +
+ + + +
+
+ {/* ===== 下方: 编辑区 ===== */} Date: Wed, 5 Aug 2026 17:57:34 +0800 Subject: [PATCH 03/48] =?UTF-8?q?fix(ai):=20=E4=BF=AE=E5=A4=8D=E6=8F=90?= =?UTF-8?q?=E4=BE=9B=E5=95=86=E5=88=87=E6=8D=A2=E4=B8=8E=20Codex=20?= =?UTF-8?q?=E8=87=AA=E5=AE=9A=E4=B9=89=E7=AB=AF=E7=82=B9?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 分离 OpenAI-compatible 与 Codex CLI 的模型和推理配置 - 复用本机 Codex provider、认证与自定义端点,并适配 Docker 回环地址 - 隔离 OpenAI 专属 reasoning_effort,修复自定义预设切换与表单保留 --- backend/app/api/settings.py | 59 ++++++++--- backend/app/services/ai_provider.py | 74 ++++++++++--- backend/tests/test_ai_provider.py | 117 +++++++++++++++++++-- frontend/src/lib/api.ts | 7 +- frontend/src/pages/settings/AI.tsx | 154 +++++++++++++++++++++------- 5 files changed, 331 insertions(+), 80 deletions(-) diff --git a/backend/app/api/settings.py b/backend/app/api/settings.py index 55f8bdf..71f8dea 100644 --- a/backend/app/api/settings.py +++ b/backend/app/api/settings.py @@ -58,7 +58,10 @@ def get_settings() -> dict: ai_configured, current_ai_model, current_codex_command, + current_codex_model, current_codex_reasoning_effort, + current_openai_model, + current_openai_reasoning_effort, ) key = secrets_store.get_tickflow_key() @@ -81,6 +84,9 @@ def get_settings() -> dict: "has_ai_key": bool(secrets_store.get_ai_key()), "ai_configured": ai_configured(ai_provider), "ai_model": current_ai_model(), + "ai_openai_model": current_openai_model(), + "ai_reasoning_effort": current_openai_reasoning_effort(), + "ai_codex_model": current_codex_model(), "ai_codex_command": current_codex_command(), "ai_codex_reasoning_effort": current_codex_reasoning_effort(), "ai_user_agent": secrets_store.get_ai_config("ai_user_agent", settings.ai_user_agent), @@ -240,6 +246,7 @@ class AiSettingsIn(BaseModel): base_url: str = "" api_key: str | None = None model: str = "" + reasoning_effort: str = Field(default="high", max_length=64) codex_command: str = "" codex_reasoning_effort: str = "" user_agent: str = "" @@ -250,12 +257,17 @@ def save_ai_settings(req: AiSettingsIn) -> dict: """保存 AI 配置(全部持久化到 secrets.json)""" from app.config import settings from app.services.ai_provider import ( + OPENAI_PROVIDER, ai_configured, current_ai_model, current_ai_provider, current_codex_command, + current_codex_model, current_codex_reasoning_effort, + current_openai_model, + current_openai_reasoning_effort, normalize_codex_command, + normalize_codex_model, normalize_codex_reasoning_effort, ) @@ -263,23 +275,8 @@ def save_ai_settings(req: AiSettingsIn) -> dict: if req.provider: updates["ai_provider"] = req.provider settings.ai_provider = req.provider - if req.base_url: - updates["ai_base_url"] = req.base_url - settings.ai_base_url = req.base_url - if req.api_key is not None: - if req.api_key: - updates["ai_api_key"] = req.api_key - settings.ai_api_key = req.api_key - else: - secrets_store.clear("ai_api_key") - settings.ai_api_key = "" - if req.provider == "codex_cli" and not req.model: - secrets_store.clear("ai_model") - settings.ai_model = "" - elif req.model: - updates["ai_model"] = req.model - settings.ai_model = req.model if req.provider == "codex_cli": + updates["ai_codex_model"] = normalize_codex_model(req.model) try: codex_command = normalize_codex_command(req.codex_command) except ValueError as exc: @@ -289,6 +286,22 @@ def save_ai_settings(req: AiSettingsIn) -> dict: updates["ai_codex_reasoning_effort"] = codex_reasoning_effort settings.ai_codex_command = codex_command settings.ai_codex_reasoning_effort = codex_reasoning_effort + else: + if req.base_url: + updates["ai_base_url"] = req.base_url + settings.ai_base_url = req.base_url + if req.api_key is not None: + if req.api_key: + updates["ai_api_key"] = req.api_key + settings.ai_api_key = req.api_key + else: + secrets_store.clear("ai_api_key") + settings.ai_api_key = "" + if req.model: + updates["ai_model"] = req.model + settings.ai_model = req.model + if req.provider == OPENAI_PROVIDER: + updates["ai_reasoning_effort"] = req.reasoning_effort.strip() # user_agent 允许清空(回到默认浏览器 UA),故无条件持久化 updates["ai_user_agent"] = req.user_agent settings.ai_user_agent = req.user_agent @@ -301,6 +314,9 @@ def save_ai_settings(req: AiSettingsIn) -> dict: "ok": True, "ai_provider": provider, "ai_model": current_ai_model(), + "ai_openai_model": current_openai_model(), + "ai_reasoning_effort": current_openai_reasoning_effort(), + "ai_codex_model": current_codex_model(), "ai_codex_command": current_codex_command(), "ai_codex_reasoning_effort": current_codex_reasoning_effort(), "ai_configured": ai_configured(provider), @@ -315,7 +331,16 @@ def clear_ai_settings() -> dict: """ from app.config import settings - secrets_store.clear("ai_provider", "ai_base_url", "ai_api_key", "ai_model", "ai_codex_command", "ai_codex_reasoning_effort") + secrets_store.clear( + "ai_provider", + "ai_base_url", + "ai_api_key", + "ai_model", + "ai_reasoning_effort", + "ai_codex_model", + "ai_codex_command", + "ai_codex_reasoning_effort", + ) # 同步重置运行时内存(provider 回默认值,其余置空) settings.ai_provider = "openai_compat" settings.ai_base_url = "" diff --git a/backend/app/services/ai_provider.py b/backend/app/services/ai_provider.py index 2886520..925d919 100644 --- a/backend/app/services/ai_provider.py +++ b/backend/app/services/ai_provider.py @@ -20,9 +20,11 @@ from app import secrets_store from app.config import settings OPENAI_COMPAT_PROVIDER = "openai_compat" +OPENAI_PROVIDER = "openai" CODEX_CLI_PROVIDER = "codex_cli" CODEX_DEFAULT_COMMAND = "codex" CODEX_SUPPORTED_REASONING_EFFORTS = {"none", "minimal", "low", "medium", "high", "xhigh"} +OPENAI_DEFAULT_REASONING_EFFORT = "high" _CODEX_ENV_ALLOWLIST = ( "PATH", @@ -104,10 +106,31 @@ def current_ai_provider() -> str: return secrets_store.get_ai_config("ai_provider", settings.ai_provider) or OPENAI_COMPAT_PROVIDER +def current_openai_model() -> str: + return secrets_store.get_ai_config("ai_model", settings.ai_model) + + +def current_codex_model() -> str: + stored = secrets_store.load() + model = stored.get("ai_codex_model") + # 旧版本的两种 provider 共用 ai_model。仅在旧配置仍启用 Codex 时回退读取, + # 避免把正常的 OpenAI-compatible 模型误当作 Codex 模型。 + if model is None and current_ai_provider() == CODEX_CLI_PROVIDER: + model = stored.get("ai_model") + return normalize_codex_model(str(model or "")) + + def current_ai_model() -> str: if current_ai_provider() == CODEX_CLI_PROVIDER: - return normalize_codex_model(str(secrets_store.load().get("ai_model") or "")) - return secrets_store.get_ai_config("ai_model", settings.ai_model) + return current_codex_model() + return current_openai_model() + + +def current_openai_reasoning_effort() -> str: + stored = secrets_store.load() + if "ai_reasoning_effort" not in stored: + return OPENAI_DEFAULT_REASONING_EFFORT + return str(stored.get("ai_reasoning_effort") or "").strip() def current_codex_command() -> str: @@ -351,10 +374,14 @@ def _is_temperature_rejected(exc: Exception) -> bool: def _openai_kwargs(*, temperature: float | None, max_tokens: int) -> dict: - """Build OpenAI create() kwargs; temperature omitted when None.""" + """Build OpenAI create() kwargs; optional parameters are omitted when empty.""" kwargs: dict = {"max_tokens": max_tokens} if temperature is not None: kwargs["temperature"] = temperature + if current_ai_provider() == OPENAI_PROVIDER: + reasoning_effort = current_openai_reasoning_effort() + if reasoning_effort: + kwargs["reasoning_effort"] = reasoning_effort return kwargs @@ -700,10 +727,14 @@ def _codex_home() -> Path: def _write_compatible_codex_config(path: Path) -> None: config = _read_codex_config() lines: list[str] = [] - local_provider = _docker_codex_local_provider(config) + active_provider = _active_codex_provider(config) - if local_provider: - lines.append(_toml_string("model_provider", "codex_local_access")) + if active_provider: + lines.append(_toml_string("model_provider", active_provider[0])) + + openai_base_url = config.get("openai_base_url") + if isinstance(openai_base_url, str) and openai_base_url: + lines.append(_toml_string("openai_base_url", openai_base_url)) model = current_ai_model() or normalize_codex_model(str(config.get("model") or "")) if model: @@ -718,41 +749,43 @@ def _write_compatible_codex_config(path: Path) -> None: lines.append(_toml_string("approval_policy", "never")) lines.append(_toml_string("sandbox_mode", "read-only")) - if local_provider: + if active_provider: + provider_name, provider = active_provider lines.append("") - lines.append("[model_providers.codex_local_access]") + lines.append(f"[model_providers.{_toml_key(provider_name)}]") for key in ("name", "base_url", "wire_api", "experimental_bearer_token"): - value = local_provider.get(key) + value = provider.get(key) if isinstance(value, str) and value: lines.append(_toml_string(key, value)) for key in ("requires_openai_auth", "supports_websockets"): - value = local_provider.get(key) + value = provider.get(key) if isinstance(value, bool): lines.append(f"{key} = {'true' if value else 'false'}") path.write_text("\n".join(lines) + "\n", encoding="utf-8") -def _docker_codex_local_provider(config: dict) -> dict | None: - """Return the local-access provider adapted to Docker's host gateway.""" - docker_host = os.environ.get("CODEX_DOCKER_HOST", "").strip() - if not docker_host or config.get("model_provider") != "codex_local_access": +def _active_codex_provider(config: dict) -> tuple[str, dict] | None: + """Return the active custom provider, adapting loopback URLs for Docker.""" + provider_name = config.get("model_provider") + if not isinstance(provider_name, str) or not provider_name: return None providers = config.get("model_providers") if not isinstance(providers, dict): return None - source = providers.get("codex_local_access") + source = providers.get(provider_name) if not isinstance(source, dict): return None provider = dict(source) base_url = str(provider.get("base_url") or "").strip() parsed = urlsplit(base_url) - if parsed.hostname in {"localhost", "127.0.0.1", "::1"}: + docker_host = os.environ.get("CODEX_DOCKER_HOST", "").strip() + if docker_host and parsed.hostname in {"localhost", "127.0.0.1", "::1"}: port = f":{parsed.port}" if parsed.port else "" provider["base_url"] = urlunsplit(parsed._replace(netloc=f"{docker_host}{port}")) - return provider + return provider_name, provider def _read_codex_config() -> dict: @@ -786,6 +819,13 @@ def _toml_string(key: str, value: str) -> str: return f'{key} = "{escaped}"' +def _toml_key(value: str) -> str: + if re.fullmatch(r"[A-Za-z0-9_-]+", value): + return value + escaped = value.replace("\\", "\\\\").replace('"', '\\"') + return f'"{escaped}"' + + def _clean_process_text(raw: bytes) -> str: text = raw.decode("utf-8", errors="replace") return _ANSI_RE.sub("", text).strip() diff --git a/backend/tests/test_ai_provider.py b/backend/tests/test_ai_provider.py index 1b5a89c..e7cf262 100644 --- a/backend/tests/test_ai_provider.py +++ b/backend/tests/test_ai_provider.py @@ -5,6 +5,9 @@ import tomllib import httpx import openai +from app import secrets_store +from app.api import settings as settings_api +from app.config import settings from app.services import ai_provider from app.services.ai_provider import ( _format_openai_error, @@ -145,6 +148,96 @@ def test_is_temperature_rejected_false_for_non_400(): assert _is_temperature_rejected(exc) is False +def test_openai_kwargs_include_configured_reasoning_effort(monkeypatch): + stored = {"ai_provider": "openai_compat"} + monkeypatch.setattr(secrets_store, "load", lambda: stored) + + assert "reasoning_effort" not in ai_provider._openai_kwargs(temperature=None, max_tokens=1000) + + stored["ai_provider"] = "openai" + assert ai_provider._openai_kwargs(temperature=None, max_tokens=1000)["reasoning_effort"] == "high" + + stored["ai_reasoning_effort"] = "custom-high" + kwargs = ai_provider._openai_kwargs(temperature=0.3, max_tokens=1000) + + assert kwargs == { + "max_tokens": 1000, + "temperature": 0.3, + "reasoning_effort": "custom-high", + } + + stored["ai_reasoning_effort"] = "" + assert "reasoning_effort" not in ai_provider._openai_kwargs(temperature=None, max_tokens=1000) + + stored["ai_reasoning_effort"] = "custom-high" + stored["ai_provider"] = "openai_compat" + assert "reasoning_effort" not in ai_provider._openai_kwargs(temperature=None, max_tokens=1000) + + +def test_ai_settings_keep_provider_models_separate(monkeypatch): + stored = { + "ai_provider": "openai_compat", + "ai_model": "custom-api-model", + } + + def save(updates: dict) -> dict: + stored.update(updates) + return stored + + def clear(*keys: str) -> dict: + for key in keys: + stored.pop(key, None) + return stored + + monkeypatch.setattr(secrets_store, "load", lambda: stored) + monkeypatch.setattr(secrets_store, "save", save) + monkeypatch.setattr(secrets_store, "clear", clear) + monkeypatch.setattr(ai_provider, "ai_configured", lambda provider=None: True) + monkeypatch.setattr(settings, "ai_provider", "openai_compat") + monkeypatch.setattr(settings, "ai_base_url", "") + monkeypatch.setattr(settings, "ai_model", "") + monkeypatch.setattr(settings, "ai_codex_command", "codex") + monkeypatch.setattr(settings, "ai_codex_reasoning_effort", "") + monkeypatch.setattr(settings, "ai_user_agent", "") + + settings_api.save_ai_settings( + settings_api.AiSettingsIn( + provider="codex_cli", + model="gpt-5.6-sol", + codex_command="codex", + codex_reasoning_effort="high", + ) + ) + + assert stored["ai_model"] == "custom-api-model" + assert stored["ai_codex_model"] == "gpt-5.6-sol" + + settings_api.save_ai_settings( + settings_api.AiSettingsIn( + provider="openai", + base_url="https://api.openai.com/v1", + model="openai-model", + reasoning_effort="vendor-high", + ) + ) + + assert stored["ai_model"] == "openai-model" + assert stored["ai_reasoning_effort"] == "vendor-high" + assert stored["ai_codex_model"] == "gpt-5.6-sol" + + settings_api.save_ai_settings( + settings_api.AiSettingsIn( + provider="openai_compat", + base_url="https://example.com/v1", + model="new-custom-model", + ) + ) + + assert stored["ai_model"] == "new-custom-model" + assert stored["ai_reasoning_effort"] == "vendor-high" + assert stored["ai_codex_model"] == "gpt-5.6-sol" + + def test_codex_process_env_excludes_application_secrets(monkeypatch, tmp_path): monkeypatch.setenv("PATH", "test-path") monkeypatch.setenv("HTTPS_PROXY", "http://proxy.example") @@ -201,7 +294,7 @@ def test_codex_config_adapts_local_access_provider_for_docker(monkeypatch, tmp_p assert provider["supports_websockets"] is False -def test_codex_config_does_not_copy_provider_without_docker_opt_in(monkeypatch, tmp_path): +def test_codex_config_preserves_remote_provider_without_docker_rewrite(monkeypatch, tmp_path): monkeypatch.delenv("CODEX_DOCKER_HOST", raising=False) monkeypatch.setattr(ai_provider, "current_ai_model", lambda: "") monkeypatch.setattr(ai_provider, "current_codex_reasoning_effort", lambda: "") @@ -209,11 +302,13 @@ def test_codex_config_does_not_copy_provider_without_docker_opt_in(monkeypatch, ai_provider, "_read_codex_config", lambda: { - "model_provider": "codex_local_access", + "model_provider": "remote-api", + "openai_base_url": "https://builtin.example/v1", "model_providers": { - "codex_local_access": { - "base_url": "http://localhost:62678/v1", - "experimental_bearer_token": "must-not-leak", + "remote-api": { + "base_url": "https://custom.example/v1", + "wire_api": "responses", + "requires_openai_auth": True, } }, }, @@ -222,7 +317,11 @@ def test_codex_config_does_not_copy_provider_without_docker_opt_in(monkeypatch, ai_provider._write_compatible_codex_config(path) - text = path.read_text(encoding="utf-8") - assert "model_provider" not in text - assert "model_providers" not in text - assert "must-not-leak" not in text + with path.open("rb") as f: + config = tomllib.load(f) + assert config["model_provider"] == "remote-api" + assert config["openai_base_url"] == "https://builtin.example/v1" + provider = config["model_providers"]["remote-api"] + assert provider["base_url"] == "https://custom.example/v1" + assert provider["wire_api"] == "responses" + assert provider["requires_openai_auth"] is True diff --git a/frontend/src/lib/api.ts b/frontend/src/lib/api.ts index c4f2e27..26d2685 100644 --- a/frontend/src/lib/api.ts +++ b/frontend/src/lib/api.ts @@ -860,6 +860,9 @@ export interface SettingsState { has_ai_key: boolean ai_configured?: boolean ai_model: string + ai_openai_model?: string + ai_reasoning_effort?: string + ai_codex_model?: string ai_codex_command?: string ai_codex_reasoning_effort?: string ai_user_agent: string @@ -1078,8 +1081,8 @@ export const api = { ), /** 保存 AI 配置 */ - saveAiSettings: (ai: { provider?: string; base_url?: string; api_key?: string; model?: string; codex_command?: string; codex_reasoning_effort?: string; user_agent?: string }) => - request<{ ok: boolean; ai_provider?: string; ai_model?: string; ai_codex_command?: string; ai_codex_reasoning_effort?: string; ai_configured?: boolean }>('/api/settings/ai', { + saveAiSettings: (ai: { provider?: string; base_url?: string; api_key?: string; model?: string; reasoning_effort?: string; codex_command?: string; codex_reasoning_effort?: string; user_agent?: string }) => + request<{ ok: boolean; ai_provider?: string; ai_model?: string; ai_openai_model?: string; ai_reasoning_effort?: string; ai_codex_model?: string; ai_codex_command?: string; ai_codex_reasoning_effort?: string; ai_configured?: boolean }>('/api/settings/ai', { method: 'POST', body: JSON.stringify(ai), }), diff --git a/frontend/src/pages/settings/AI.tsx b/frontend/src/pages/settings/AI.tsx index ed7b9fb..76f302d 100644 --- a/frontend/src/pages/settings/AI.tsx +++ b/frontend/src/pages/settings/AI.tsx @@ -1,4 +1,4 @@ -import { useState, useEffect } from 'react' +import { useState, useEffect, useRef } from 'react' import { useMutation, useQueryClient } from '@tanstack/react-query' import { Save, Loader2, Check, Wifi, WifiOff, Eye, EyeOff, Shield, @@ -14,10 +14,13 @@ const INPUT_CLS = 'w-full h-9 px-2.5 rounded-lg bg-base border-0 ring-1 ring-border/30 text-xs font-mono text-foreground placeholder:text-muted/30 focus:outline-none focus:ring-2 focus:ring-accent/30 transition-shadow' const CODEX_PROVIDER = 'codex_cli' -const OPENAI_PROVIDER = 'openai_compat' +const OPENAI_PROVIDER = 'openai' +const OPENAI_COMPAT_PROVIDER = 'openai_compat' const CODEX_COMMAND = 'codex' const DEFAULT_CODEX_MODEL = 'gpt-5.6-sol' const DEFAULT_CODEX_REASONING_EFFORT = 'xhigh' +const DEFAULT_OPENAI_MODEL = 'gpt-5.5' +const DEFAULT_REASONING_EFFORT = 'high' const SAVED_CODEX_OPTION_VALUE = '__saved_codex_config__' const CODEX_REASONING_LABELS: Record = { high: '高', @@ -42,8 +45,11 @@ const codexModelLabel = (model?: string, effort?: string) => { return effortLabel ? `${modelLabel} · ${effortLabel}` : modelLabel } -const PRESETS: { label: string; provider?: string; url: string; model: string; codexCommand?: string; website: string; websiteLabel: string; description: string; custom?: boolean }[] = [ +type AiPreset = { label: string; provider?: string; url: string; model: string; codexCommand?: string; website: string; websiteLabel: string; description: string; custom?: boolean } + +const PRESETS: AiPreset[] = [ { label: '自定义', url: '', model: '', website: '', websiteLabel: '', description: '不自动填充任何配置,完全手动填写 API 地址、模型和密钥。', custom: true }, + { label: 'OpenAI', provider: OPENAI_PROVIDER, url: 'https://api.openai.com/v1', model: DEFAULT_OPENAI_MODEL, website: 'https://platform.openai.com/', websiteLabel: 'platform.openai.com', description: 'OpenAI 官方接口,可单独配置模型支持的推理强度。' }, { label: 'DeepSeek', url: 'https://api.deepseek.com', model: 'deepseek-v4-pro', website: 'https://www.deepseek.com/', websiteLabel: 'deepseek.com', description: 'DeepSeek 官方 OpenAI 兼容接口。' }, { label: '通义千问', url: 'https://dashscope.aliyuncs.com/compatible-mode/v1', model: 'qwen-3.6plus', website: 'https://tongyi.aliyun.com/', websiteLabel: 'tongyi.aliyun.com', description: '阿里云 DashScope 兼容模式接口。' }, { label: '智谱 GLM', url: 'https://open.bigmodel.cn/api/paas/v4', model: 'glm-5.2', website: 'https://open.bigmodel.cn/', websiteLabel: 'open.bigmodel.cn', description: '智谱 AI 官方 OpenAI 兼容接口。' }, @@ -52,15 +58,23 @@ const PRESETS: { label: string; provider?: string; url: string; model: string; c { label: '炸鸡中转站', url: 'https://api.zhaji.dev/v1', model: 'gpt-5.5', website: 'https://api.zhaji.dev', websiteLabel: 'api.zhaji.dev', description: 'OpenAI 兼容中转服务,适合直接使用国际模型。' }, ] +const findPreset = (provider: string, baseUrl: string, codexCommand: string) => PRESETS.find(p => { + if (p.custom || (p.provider ?? OPENAI_COMPAT_PROVIDER) !== provider) return false + if (provider === OPENAI_PROVIDER) return true + return provider === CODEX_PROVIDER ? p.codexCommand === codexCommand : p.url === baseUrl +}) ?? PRESETS[0] + export function SettingsAIPanel() { const qc = useQueryClient() const settings = useSettings() const s = settings.data - const [provider, setProvider] = useState(OPENAI_PROVIDER) + const [provider, setProvider] = useState(OPENAI_COMPAT_PROVIDER) const [baseUrl, setBaseUrl] = useState('') const [apiKey, setApiKey] = useState('') const [model, setModel] = useState('') + const [reasoningEffort, setReasoningEffort] = useState(DEFAULT_REASONING_EFFORT) + const [codexModel, setCodexModel] = useState('') const [codexReasoningEffort, setCodexReasoningEffort] = useState('') const [codexCommand, setCodexCommand] = useState(CODEX_COMMAND) const [customUa, setCustomUa] = useState(false) @@ -70,15 +84,21 @@ export function SettingsAIPanel() { const [confirmClear, setConfirmClear] = useState(false) const [testing, setTesting] = useState(false) const [testResult, setTestResult] = useState<{ ok: boolean; msg: string } | null>(null) + const [selectedPresetLabel, setSelectedPresetLabel] = useState(PRESETS[0].label) + const directDrafts = useRef({ + custom: { baseUrl: '', model: '' }, + openai: { baseUrl: 'https://api.openai.com/v1', model: DEFAULT_OPENAI_MODEL }, + }) + const draftsInitialized = useRef(false) const isCodexProvider = provider === CODEX_PROVIDER + const isOpenAIProvider = provider === OPENAI_PROVIDER const savedCodexProvider = s?.ai_provider === CODEX_PROVIDER const configured = s?.ai_configured ?? (savedCodexProvider ? !!(s?.ai_codex_command ?? CODEX_COMMAND) : s?.has_ai_key) - // 选中的预设: 精确匹配 provider+url/codexCommand; 匹配不上时默认"自定义" - const matchedPreset = PRESETS.find(p => (p.provider ?? OPENAI_PROVIDER) === provider && (isCodexProvider ? p.codexCommand === codexCommand : p.url === baseUrl)) - const selectedPreset = matchedPreset ?? PRESETS.find(p => p.custom) - const savedCodexModel = savedCodexProvider ? (s?.ai_model ?? '') : '' - const savedCodexEffort = savedCodexProvider ? (s?.ai_codex_reasoning_effort ?? '') : '' + const selectedPreset = PRESETS.find(p => p.label === selectedPresetLabel) ?? PRESETS[0] + const configTitle = isCodexProvider ? 'Codex CLI 配置' : isOpenAIProvider ? 'OpenAI 配置' : selectedPreset.custom ? '自定义配置' : `${selectedPreset.label} 配置` + const savedCodexModel = s?.ai_codex_model ?? (savedCodexProvider ? (s?.ai_model ?? '') : '') + const savedCodexEffort = s?.ai_codex_reasoning_effort ?? '' const savedCodexOptionKnown = CODEX_MODEL_OPTIONS.some(option => option.model === savedCodexModel && option.effort === savedCodexEffort, ) @@ -96,7 +116,7 @@ export function SettingsAIPanel() { ? [savedCodexOption, ...CODEX_MODEL_OPTIONS] : CODEX_MODEL_OPTIONS const selectedCodexModelOption = codexModelOptions.find(option => - option.model === model && option.effort === codexReasoningEffort, + option.model === codexModel && option.effort === codexReasoningEffort, ) ?? CODEX_MODEL_OPTIONS[0] const codexModelSelectValue = selectedCodexModelOption.value const canSave = isCodexProvider ? true : !!baseUrl.trim() && !!model.trim() @@ -105,11 +125,26 @@ export function SettingsAIPanel() { if (!s) return // 未配置过 AI (无 api_key): 字段留空, 默认选中"自定义"预设, 不预填充后端默认值 const unconfigured = !s.has_ai_key && !s.ai_configured - const savedProvider = s.ai_provider ?? OPENAI_PROVIDER + const savedProvider = s.ai_provider ?? OPENAI_COMPAT_PROVIDER + const savedBaseUrl = unconfigured ? '' : (s.ai_base_url ?? '') + const savedOpenAIModel = unconfigured ? '' : (s.ai_openai_model ?? (savedProvider !== CODEX_PROVIDER ? s.ai_model : '') ?? '') + const savedPreset = unconfigured ? PRESETS[0] : findPreset(savedProvider, savedBaseUrl, s.ai_codex_command ?? CODEX_COMMAND) + if (!draftsInitialized.current) { + const officialOpenAI = PRESETS.find(p => p.provider === OPENAI_PROVIDER) + if (savedProvider === OPENAI_PROVIDER || (savedProvider === CODEX_PROVIDER && savedBaseUrl === officialOpenAI?.url)) { + directDrafts.current.openai = { baseUrl: savedBaseUrl, model: savedOpenAIModel } + } else if (findPreset(OPENAI_COMPAT_PROVIDER, savedBaseUrl, CODEX_COMMAND).custom) { + directDrafts.current.custom = { baseUrl: savedBaseUrl, model: savedOpenAIModel } + } + draftsInitialized.current = true + } setProvider(savedProvider) - setBaseUrl(unconfigured ? '' : (s.ai_base_url ?? '')) - setModel(unconfigured ? '' : (s.ai_model ?? '')) - setCodexReasoningEffort(unconfigured ? '' : (s.ai_codex_reasoning_effort ?? '')) + setSelectedPresetLabel(savedPreset.label) + setBaseUrl(savedBaseUrl) + setModel(savedOpenAIModel) + setReasoningEffort(s.ai_reasoning_effort ?? DEFAULT_REASONING_EFFORT) + setCodexModel(s.ai_codex_model ?? (savedProvider === CODEX_PROVIDER ? s.ai_model : '') ?? '') + setCodexReasoningEffort(s.ai_codex_reasoning_effort ?? '') setCodexCommand(s.ai_codex_command ?? CODEX_COMMAND) const ua = s.ai_user_agent ?? '' setCustomUa(!!ua) @@ -120,7 +155,8 @@ export function SettingsAIPanel() { provider, base_url: baseUrl, api_key: apiKey || undefined, - model, + model: isCodexProvider ? codexModel : model, + ...(isOpenAIProvider ? { reasoning_effort: reasoningEffort } : {}), codex_command: isCodexProvider ? CODEX_COMMAND : codexCommand, codex_reasoning_effort: isCodexProvider ? codexReasoningEffort : '', user_agent: customUa ? userAgent : '', @@ -135,7 +171,10 @@ export function SettingsAIPanel() { ...prev, ai_provider: result.ai_provider ?? provider, ai_base_url: baseUrl, - ai_model: result.ai_model ?? model, + ai_model: result.ai_model ?? (isCodexProvider ? codexModel : model), + ai_openai_model: result.ai_openai_model ?? model, + ai_reasoning_effort: result.ai_reasoning_effort ?? reasoningEffort, + ai_codex_model: result.ai_codex_model ?? codexModel, ai_codex_command: result.ai_codex_command ?? (isCodexProvider ? CODEX_COMMAND : codexCommand), ai_codex_reasoning_effort: result.ai_codex_reasoning_effort ?? (isCodexProvider ? codexReasoningEffort : ''), ai_configured: result.ai_configured ?? (isCodexProvider ? true : (apiKey ? true : prev.ai_configured)), @@ -153,18 +192,28 @@ export function SettingsAIPanel() { mutationFn: () => api.clearAiSettings(), onSuccess: () => { setConfirmClear(false) - setProvider(OPENAI_PROVIDER) + setProvider(OPENAI_COMPAT_PROVIDER) + setSelectedPresetLabel(PRESETS[0].label) setBaseUrl('') setApiKey('') setModel('') + setReasoningEffort(DEFAULT_REASONING_EFFORT) + setCodexModel('') setCodexReasoningEffort('') setCodexCommand(CODEX_COMMAND) + directDrafts.current = { + custom: { baseUrl: '', model: '' }, + openai: { baseUrl: 'https://api.openai.com/v1', model: DEFAULT_OPENAI_MODEL }, + } setTestResult(null) qc.setQueryData(QK.settings, prev => prev ? { ...prev, - ai_provider: OPENAI_PROVIDER, + ai_provider: OPENAI_COMPAT_PROVIDER, ai_base_url: '', ai_model: '', + ai_openai_model: '', + ai_reasoning_effort: DEFAULT_REASONING_EFFORT, + ai_codex_model: '', ai_codex_command: CODEX_COMMAND, ai_codex_reasoning_effort: '', has_ai_key: false, @@ -186,20 +235,42 @@ export function SettingsAIPanel() { setUserAgent(`Mozilla/5.0 (${pf}) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/${major}.0.0.0 Safari/537.36`) } - const handlePreset = (p: typeof PRESETS[number]) => { + const handlePreset = (p: AiPreset) => { + setSelectedPresetLabel(p.label) if (p.custom) { - // 自定义: 清空所有自动填充字段, 由用户完全手动填写 - setProvider(OPENAI_PROVIDER) - setBaseUrl('') - setModel('') - setCodexReasoningEffort('') + setProvider(OPENAI_COMPAT_PROVIDER) + setBaseUrl(directDrafts.current.custom.baseUrl) + setModel(directDrafts.current.custom.model) return } - setProvider(p.provider ?? OPENAI_PROVIDER) - setBaseUrl(p.url) - setModel(p.model) - setCodexReasoningEffort(p.provider === CODEX_PROVIDER ? DEFAULT_CODEX_REASONING_EFFORT : '') - if (p.codexCommand) setCodexCommand(CODEX_COMMAND) + if (p.provider === CODEX_PROVIDER) { + setProvider(CODEX_PROVIDER) + setCodexModel(p.model) + setCodexReasoningEffort(DEFAULT_CODEX_REASONING_EFFORT) + setCodexCommand(CODEX_COMMAND) + return + } + const nextProvider = p.provider ?? OPENAI_COMPAT_PROVIDER + setProvider(nextProvider) + if (nextProvider === OPENAI_PROVIDER) { + setBaseUrl(directDrafts.current.openai.baseUrl) + setModel(directDrafts.current.openai.model) + } else { + setBaseUrl(p.url) + setModel(p.model) + } + } + + const handleBaseUrlChange = (value: string) => { + setBaseUrl(value) + if (selectedPreset.custom) directDrafts.current.custom.baseUrl = value + if (isOpenAIProvider) directDrafts.current.openai.baseUrl = value + } + + const handleModelChange = (value: string) => { + setModel(value) + if (selectedPreset.custom) directDrafts.current.custom.model = value + if (isOpenAIProvider) directDrafts.current.openai.model = value } const handleTest = async () => { @@ -280,7 +351,7 @@ export function SettingsAIPanel() { {isCodexProvider ? 'codex exec' : 'Chat Completions'} @@ -290,7 +361,7 @@ export function SettingsAIPanel() { >
{isCodexProvider ? ( -
+
{CODEX_COMMAND} @@ -305,7 +376,7 @@ export function SettingsAIPanel() { onChange={e => { const value = e.target.value const option = codexModelOptions.find(item => item.value === value) ?? CODEX_MODEL_OPTIONS[0] - setModel(option.model) + setCodexModel(option.model) setCodexReasoningEffort(option.effort) }} className={INPUT_CLS} @@ -318,15 +389,28 @@ export function SettingsAIPanel() {
) : ( <> -
+
- setBaseUrl(e.target.value)} placeholder="https://api.zhaji.dev/v1" className={INPUT_CLS} /> + handleBaseUrlChange(e.target.value)} placeholder="https://api.zhaji.dev/v1" className={INPUT_CLS} /> - setModel(e.target.value)} placeholder="gpt-5.6-sol" className={INPUT_CLS} /> + handleModelChange(e.target.value)} placeholder="gpt-5.6-sol" className={INPUT_CLS} />
+ {isOpenAIProvider && ( +
+
+ OpenAI 专属 +
+
+ + setReasoningEffort(e.target.value)} placeholder={DEFAULT_REASONING_EFFORT} className={INPUT_CLS} /> + +
+
+ )} +
From e1c6aa2a14d0bf8cf4c85040b157fe7767142018 Mon Sep 17 00:00:00 2001 From: Arepeater Date: Sat, 8 Aug 2026 15:21:26 +0800 Subject: [PATCH 04/48] =?UTF-8?q?fix(config):=20=E6=94=B6=E6=95=9B=20Vite?= =?UTF-8?q?=20=E9=85=8D=E7=BD=AE=E5=B9=B6=E5=85=BC=E5=AE=B9=20AI=20?= =?UTF-8?q?=E5=8F=AF=E9=80=89=E5=8F=82=E6=95=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 移除重复的 Vite JS 配置并阻止 TypeScript 构建重新生成 - 精确识别 temperature 与 reasoning_effort 的 400 错误并有界降级 - 补充 Docker 只读挂载 .env 的安全边界说明 --- backend/app/services/ai_provider.py | 91 +++++++++++++++++++---------- backend/tests/test_ai_provider.py | 25 +++++++- docs/configuration.md | 2 +- frontend/tsconfig.node.json | 3 +- frontend/vite.config.js | 54 ----------------- 5 files changed, 87 insertions(+), 88 deletions(-) delete mode 100644 frontend/vite.config.js diff --git a/backend/app/services/ai_provider.py b/backend/app/services/ai_provider.py index 925d919..754a485 100644 --- a/backend/app/services/ai_provider.py +++ b/backend/app/services/ai_provider.py @@ -269,23 +269,20 @@ async def _run_openai_once( client = _openai_client(ai_key, timeout) model = current_ai_model() req_messages = list(messages) - try: - resp = await client.chat.completions.create( - model=model, - messages=req_messages, - **_openai_kwargs(temperature=temperature, max_tokens=max_tokens), - ) - except Exception as exc: - # Reasoning 类模型 (如 kimi-k2.7-code, deepseek-r1, o 系列) 拒绝非约定 - # temperature (Moonshot 报 "only 1 is allowed for this model")。不再靠 - # 模型名猜测, 而是捕获该错误后去掉 temperature 重试一次 —— 对所有此类模型都稳。 - if temperature is not None and _is_temperature_rejected(exc): + kwargs = _openai_kwargs(temperature=temperature, max_tokens=max_tokens) + while True: + try: resp = await client.chat.completions.create( model=model, messages=req_messages, - **_openai_kwargs(temperature=None, max_tokens=max_tokens), + **kwargs, ) - else: + break + except Exception as exc: + retry_kwargs = _openai_retry_kwargs(exc, kwargs) + if retry_kwargs is not None: + kwargs = retry_kwargs + continue if _is_openai_transport_error(exc): raise RuntimeError(_format_openai_error(exc)) from exc raise @@ -315,23 +312,22 @@ async def _stream_openai( if delta and delta.content: yield delta.content - try: - stream = await client.chat.completions.create( - model=model, - messages=req_messages, - **_openai_kwargs(temperature=temperature, max_tokens=max_tokens), - stream=True, - ) - except Exception as exc: - # 流尚未开始 yield, 可安全重建: 去掉 temperature 后重开 stream。 - if temperature is not None and _is_temperature_rejected(exc): + kwargs = _openai_kwargs(temperature=temperature, max_tokens=max_tokens) + while True: + try: stream = await client.chat.completions.create( model=model, messages=req_messages, - **_openai_kwargs(temperature=None, max_tokens=max_tokens), + **kwargs, stream=True, ) - else: + break + except Exception as exc: + # 流尚未开始 yield, 可安全移除被拒绝的可选参数后重建。 + retry_kwargs = _openai_retry_kwargs(exc, kwargs) + if retry_kwargs is not None: + kwargs = retry_kwargs + continue if _is_openai_transport_error(exc): raise RuntimeError(_format_openai_error(exc)) from exc raise @@ -358,11 +354,10 @@ def _openai_client(api_key: str, timeout: float): ) -# Reasoning / thinking 类模型 (kimi-k2.7-code, deepseek-r1, OpenAI o 系列等) 不接受 -# 任意 temperature, 上游会以 400 拒绝 (如 Moonshot: "only 1 is allowed for this model")。 -# 这里不靠模型名猜测, 而是在真正命中该错误后自动去掉 temperature 重试 (见 -# _run_openai_once / _stream_openai), 对任意 reasoning 模型都稳健。 -_TEMP_REJECT_HINTS = ("temperature", "only 1 is allowed", "unsupported parameter") +# 不同模型可能拒绝 temperature 或 reasoning_effort。这里不靠模型名猜测, +# 只在 400 明确指出对应参数时移除该参数并重试; 每个参数最多移除一次。 +_TEMP_REJECT_HINTS = ("temperature", "only 1 is allowed") +_REASONING_EFFORT_REJECT_HINTS = ("reasoning_effort", "reasoning effort") def _is_temperature_rejected(exc: Exception) -> bool: @@ -370,7 +365,41 @@ def _is_temperature_rejected(exc: Exception) -> bool: if getattr(exc, "status_code", None) != 400: return False text = _openai_error_detail(exc) or str(exc) - return any(h in text.lower() for h in _TEMP_REJECT_HINTS) + return _openai_error_param(exc) == "temperature" or any( + h in text.lower() for h in _TEMP_REJECT_HINTS + ) + + +def _is_reasoning_effort_rejected(exc: Exception) -> bool: + """True if the upstream 400 specifically rejects reasoning_effort.""" + if getattr(exc, "status_code", None) != 400: + return False + text = _openai_error_detail(exc) or str(exc) + return _openai_error_param(exc) == "reasoning_effort" or any( + h in text.lower() for h in _REASONING_EFFORT_REJECT_HINTS + ) + + +def _openai_error_param(exc: Exception) -> str: + body = getattr(exc, "body", None) + if not isinstance(body, dict): + return "" + error = body.get("error") + if isinstance(error, dict): + body = error + return str(body.get("param") or "").strip().lower() + + +def _openai_retry_kwargs(exc: Exception, kwargs: dict) -> dict | None: + """Remove one explicitly rejected optional argument for a bounded retry.""" + retry_kwargs = dict(kwargs) + if "temperature" in retry_kwargs and _is_temperature_rejected(exc): + retry_kwargs.pop("temperature") + return retry_kwargs + if "reasoning_effort" in retry_kwargs and _is_reasoning_effort_rejected(exc): + retry_kwargs.pop("reasoning_effort") + return retry_kwargs + return None def _openai_kwargs(*, temperature: float | None, max_tokens: int) -> dict: diff --git a/backend/tests/test_ai_provider.py b/backend/tests/test_ai_provider.py index e7cf262..75930fb 100644 --- a/backend/tests/test_ai_provider.py +++ b/backend/tests/test_ai_provider.py @@ -111,7 +111,7 @@ def test_is_temperature_rejected_matches_moonshot_message(): assert _is_temperature_rejected(exc) is True -def test_is_temperature_rejected_matches_generic_temperature_hint(): +def test_optional_openai_params_use_targeted_400_fallbacks(): response = httpx.Response( 400, json={"error": {"message": "unsupported parameter: temperature"}}, @@ -123,6 +123,29 @@ def test_is_temperature_rejected_matches_generic_temperature_hint(): ) assert _is_temperature_rejected(exc) is True + kwargs = {"max_tokens": 1000, "temperature": 0.3, "reasoning_effort": "high"} + assert ai_provider._openai_retry_kwargs(exc, kwargs) == { + "max_tokens": 1000, + "reasoning_effort": "high", + } + + response = httpx.Response( + 400, + json={"error": {"message": "unrecognized request argument", "param": "reasoning_effort"}}, + request=httpx.Request("POST", "https://example.com/v1/chat/completions"), + ) + exc = openai.BadRequestError( + "bad request", response=response, + body={"error": {"message": "unrecognized request argument", "param": "reasoning_effort"}}, + ) + assert _is_temperature_rejected(exc) is False + assert ai_provider._is_reasoning_effort_rejected(exc) is True + assert ai_provider._openai_retry_kwargs(exc, kwargs) == { + "max_tokens": 1000, + "temperature": 0.3, + } + assert kwargs == {"max_tokens": 1000, "temperature": 0.3, "reasoning_effort": "high"} + def test_is_temperature_rejected_false_for_other_400(): """非 temperature 相关的 400 (如 model not found) 不应触发去 temperature 重试。""" diff --git a/docs/configuration.md b/docs/configuration.md index 21c38ad..65546a8 100644 --- a/docs/configuration.md +++ b/docs/configuration.md @@ -88,7 +88,7 @@ AUTH_PASSWORD='你的密码' # 至少 6 位;仅首次生效,已设过则不覆 ``` 面板首次设置访问密码时,出于安全考虑**仅允许本机或内网访问**(防公网陌生人抢先设置锁死面板)。公网服务器部署可通过此环境变量预置首个密码。 -密码建议使用单引号包裹,避免 Docker Compose 插值 `$VAR`;Docker 启动时也会只读挂载原始 `.env`,兼容已有的未加引号配置。 +密码建议使用单引号包裹,Docker 启动时会把整个原始 `.env` 只读挂载到容器内 `/app/.env`,兼容已有的未加引号配置。容器可以读取其中的密钥但不能修改该文件,请保持主机文件权限为 `600` 并仅运行可信镜像。 详细步骤、SSH 转发方案、重置密码方法见 [deployment.md → 访问密码设置](./deployment.md#访问密码设置公网部署必读)。 diff --git a/frontend/tsconfig.node.json b/frontend/tsconfig.node.json index ba173d2..870aa8d 100644 --- a/frontend/tsconfig.node.json +++ b/frontend/tsconfig.node.json @@ -8,7 +8,8 @@ "moduleResolution": "bundler", "allowSyntheticDefaultImports": true, "strict": true, - "composite": true + "composite": true, + "emitDeclarationOnly": true }, "include": ["vite.config.ts"] } diff --git a/frontend/vite.config.js b/frontend/vite.config.js deleted file mode 100644 index 239f479..0000000 --- a/frontend/vite.config.js +++ /dev/null @@ -1,54 +0,0 @@ -import { defineConfig } from 'vite'; -import react from '@vitejs/plugin-react'; -import path from 'node:path'; -const backendHost = process.env.BACKEND_HOST || '127.0.0.1'; -const proxyHost = ['0.0.0.0', '::'].includes(backendHost) ? '127.0.0.1' : backendHost; -const backendPort = process.env.BACKEND_PORT || '3018'; -const backendTarget = `http://${proxyHost}:${backendPort}`; -export default defineConfig({ - plugins: [react()], - resolve: { - alias: { - '@': path.resolve(__dirname, './src'), - }, - }, - server: { - host: '0.0.0.0', // dev.sh / dev.ps1 会用 CLI --host 覆盖 - port: 3011, - proxy: { - // dev 时 /api 转发到与启动脚本相同的 FastAPI 地址 - '/api': { - target: backendTarget, - // SSE 端点需要禁用缓冲 - configure: (proxy) => { - proxy.on('proxyReq', (_proxyReq, req) => { - if (req.url?.includes('/stream')) { - _proxyReq.setHeader('Accept', 'text/event-stream'); - _proxyReq.setHeader('Cache-Control', 'no-cache'); - _proxyReq.setHeader('Connection', 'keep-alive'); - } - }); - }, - }, - '/health': backendTarget, - }, - }, - build: { - outDir: 'dist', - sourcemap: false, - rollupOptions: { - output: { - // 把重型图表库拆到独立 chunk, 避免打进主包 + 让页面按需加载。 - // 用函数形式按 node_modules 路径匹配, 比对象形式更可靠。 - manualChunks(id) { - if (id.includes('node_modules')) { - if (id.includes('echarts')) - return 'echarts'; - if (id.includes('lightweight-charts')) - return 'lightweight-charts'; - } - }, - }, - }, - }, -}); From 2a2d3b764e7c16178b29cad025811ffa4f6aa5b2 Mon Sep 17 00:00:00 2001 From: shy3130 <415333856@qq.com> Date: Mon, 10 Aug 2026 11:39:08 +0800 Subject: [PATCH 05/48] feat: add watchlist groups and multi-day intraday view --- backend/app/api/kline.py | 129 +++++- backend/app/api/watchlist.py | 74 +++- backend/app/services/watchlist.py | 291 +++++++++++--- backend/tests/test_minute_range_api.py | 107 +++++ backend/tests/test_watchlist_groups.py | 114 ++++++ backend/uv.lock | 13 +- frontend/src/components/EChartsIntraday.tsx | 50 +-- .../components/EChartsMultiDayIntraday.tsx | 373 ++++++++++++++++++ frontend/src/components/StockInfoBar.tsx | 44 ++- .../components/StockMultiDayIntradayChart.tsx | 152 +++++++ frontend/src/components/StockPanel.tsx | 14 +- .../src/components/StockPreviewDialog.tsx | 173 +++++--- frontend/src/components/WatchlistAddMenu.tsx | 268 +++++++++++++ frontend/src/components/WatchlistGroups.tsx | 359 +++++++++++++++++ .../src/components/WatchlistImportDialog.tsx | 31 +- .../src/components/screener/ScreenerTable.tsx | 41 +- frontend/src/lib/api.ts | 73 +++- frontend/src/lib/intraday-chart.ts | 40 ++ frontend/src/lib/queryKeys.ts | 3 + frontend/src/lib/storage.ts | 3 + frontend/src/lib/useSharedMutations.ts | 10 +- frontend/src/lib/watchlist-group-colors.ts | 33 ++ frontend/src/pages/Screener.tsx | 32 +- frontend/src/pages/Watchlist.tsx | 225 +++++++++-- .../src/pages/backtest/StrategyBacktest.tsx | 38 +- 25 files changed, 2433 insertions(+), 257 deletions(-) create mode 100644 backend/tests/test_minute_range_api.py create mode 100644 backend/tests/test_watchlist_groups.py create mode 100644 frontend/src/components/EChartsMultiDayIntraday.tsx create mode 100644 frontend/src/components/StockMultiDayIntradayChart.tsx create mode 100644 frontend/src/components/WatchlistAddMenu.tsx create mode 100644 frontend/src/components/WatchlistGroups.tsx create mode 100644 frontend/src/lib/intraday-chart.ts create mode 100644 frontend/src/lib/watchlist-group-colors.ts diff --git a/backend/app/api/kline.py b/backend/app/api/kline.py index 5dcde88..14744f1 100644 --- a/backend/app/api/kline.py +++ b/backend/app/api/kline.py @@ -268,6 +268,47 @@ def _get_price_limit_info( return info +def _get_previous_closes( + repo, + symbol: str, + trade_dates: list[date], + asset_type: str, +) -> dict[date, float | None]: + """Return the previous trading day's adjusted close for each session.""" + if not trade_dates: + return {} + start = min(trade_dates) - timedelta(days=45) + end = max(trade_dates) + try: + daily = repo.get_daily_asset( + asset_type, + symbol, + start, + end, + columns=["date", "close"], + ).sort("date") + except Exception: + daily = None + if daily is None or daily.is_empty(): + return {trade_date: None for trade_date in trade_dates} + + closes: list[tuple[date, float]] = [] + for daily_date, close in daily.select(["date", "close"]).iter_rows(): + if close is None: + continue + numeric = float(close) + if math.isfinite(numeric) and numeric > 0: + closes.append((daily_date, numeric)) + + result: dict[date, float | None] = {} + for trade_date in trade_dates: + result[trade_date] = next( + (close for daily_date, close in reversed(closes) if daily_date < trade_date), + None, + ) + return result + + @router.get("/daily") def get_daily( request: Request, @@ -679,6 +720,73 @@ def get_minute_batch(request: Request, body: dict): return {"data": result} +@router.get("/minute-range") +def get_minute_range( + request: Request, + symbol: str = Query(..., description="标的代码"), + days: int = Query(10, ge=1, le=20, description="最近交易日数量"), +): + """读取单只标的最近 N 个已落库交易日的分钟 K。""" + import polars as pl + + repo = request.app.state.repo + asset_type = repo.resolve_asset_type(symbol) + stock_info = ( + _get_stock_info(repo, symbol) + if asset_type == "stock" + else _get_asset_info(repo, symbol, asset_type) + ) + base_response = { + "symbol": symbol, + "name": stock_info.get("name"), + "asset_type": asset_type, + "requested_days": days, + } + + # 指数分钟 K 不落本地仓库, 最新分时仍由 /api/index/minute 实时读取。 + if asset_type == "index": + return {**base_response, "sessions": [], "source": "none"} + + end = cn_today() + start = end - timedelta(days=days * 3 + 20) + minute = repo.get_minute_range([symbol], start, end, asset_type=asset_type) + if minute.is_empty() or "datetime" not in minute.columns: + return {**base_response, "sessions": [], "source": "none"} + + minute = minute.with_columns( + pl.col("datetime").dt.date().alias("_trade_date"), + ) + trade_dates = sorted(minute["_trade_date"].unique().to_list())[-days:] + previous_closes = _get_previous_closes(repo, symbol, trade_dates, asset_type) + row_columns = [ + column + for column in ( + "datetime", "open", "high", "low", "close", "volume", "amount" + ) + if column in minute.columns + ] + sessions = [] + for trade_date in trade_dates: + rows = ( + minute.filter(pl.col("_trade_date") == trade_date) + .sort("datetime") + .select(row_columns) + .to_dicts() + ) + if rows: + sessions.append({ + "date": trade_date.isoformat(), + "prev_close": previous_closes.get(trade_date), + "rows": rows, + }) + + return { + **base_response, + "sessions": sessions, + "source": "local" if sessions else "none", + } + + @router.get("/minute") def get_minute( request: Request, @@ -721,13 +829,20 @@ def get_minute( price_limit = _get_price_limit_info( repo, symbol, trade_date, asset_type, stock_name, ) + prev_close = _get_previous_closes( + repo, symbol, [trade_date], asset_type, + ).get(trade_date) return { "symbol": symbol, "name": stock_name, "stock_info": stock_info, "date": str(trade_date), "rows": df.to_dicts(), "source": "live", "asset_type": asset_type, "price_limit": price_limit, + "prev_close": prev_close, } + prev_close = _get_previous_closes( + repo, symbol, [trade_date], asset_type, + ).get(trade_date) price_limit = _get_price_limit_info( repo, symbol, trade_date, asset_type, stock_name, ) @@ -758,6 +873,7 @@ def get_minute( "date": str(trade_date), "rows": df.to_dicts(), "source": "local", "asset_type": asset_type, "price_limit": price_limit, + "prev_close": prev_close, } # 本地不完整或无数据 → 从 TickFlow 实时拉取 @@ -768,6 +884,7 @@ def get_minute( "source": "live" if not live_df.is_empty() else "none", "asset_type": asset_type, "price_limit": price_limit, + "prev_close": prev_close, } @@ -909,13 +1026,21 @@ async def sync_minute_single(request: Request, body: dict): body: { "symbol": "000001.SZ" } 用于个股分时图"获取数据"按钮: 本地无数据时单独拉取并持久化。 """ + import asyncio + from app.services.preferences import get_minute_sync_days - from app.tickflow.capabilities import Cap symbol = body.get("symbol", "").strip() if not symbol: raise HTTPException(status_code=400, detail="symbol 不能为空") + requested_days = body.get("days") + if requested_days is not None: + if isinstance(requested_days, bool) or not isinstance(requested_days, int): + raise HTTPException(status_code=400, detail="days 必须是整数") + if requested_days < 1 or requested_days > 30: + raise HTTPException(status_code=400, detail="days 必须在 1 到 30 之间") + repo = request.app.state.repo capset = request.app.state.capabilities @@ -927,7 +1052,7 @@ async def sync_minute_single(request: Request, body: dict): if not _minute_allowed(capset): raise HTTPException(status_code=403, detail="需要 Pro+ 权限") - days = get_minute_sync_days() + days = requested_days if requested_days is not None else get_minute_sync_days() loop = asyncio.get_event_loop() def _run(): diff --git a/backend/app/api/watchlist.py b/backend/app/api/watchlist.py index 8780b16..867bf53 100644 --- a/backend/app/api/watchlist.py +++ b/backend/app/api/watchlist.py @@ -36,11 +36,22 @@ _OCR_LIMITER = anyio.CapacityLimiter(2) class AddRequest(BaseModel): symbol: str note: str = "" + group_id: str | None = None class BatchAddRequest(BaseModel): symbols: list[str] note: str = "" + group_id: str | None = None + + +class GroupNameRequest(BaseModel): + name: str + color: str | None = None + + +class GroupAssignRequest(BaseModel): + group_id: str | None = None def _with_names(rows: list[dict], request: Request) -> list[dict]: @@ -64,20 +75,54 @@ def list_all(request: Request): @router.post("") def add_one(req: AddRequest, request: Request): - rows = watchlist.add(req.symbol, req.note) + try: + rows = watchlist.add(req.symbol, req.note, req.group_id) + except ValueError as e: + raise HTTPException(400, str(e)) from e return {"symbols": _with_names(rows, request)} @router.post("/batch") def add_batch(req: BatchAddRequest, request: Request): - existing = {r["symbol"] for r in watchlist.list_symbols()} - added = 0 - for sym in req.symbols: - if sym not in existing: - added += 1 - existing.add(sym) - watchlist.add(sym, req.note) - return {"symbols": _with_names(watchlist.list_symbols(), request), "added": added} + try: + rows, added = watchlist.add_batch(req.symbols, req.note, req.group_id) + except ValueError as e: + raise HTTPException(400, str(e)) from e + return {"symbols": _with_names(rows, request), "added": added} + + +@router.get("/groups") +def list_groups(): + return {"groups": watchlist.list_groups()} + + +@router.post("/groups") +def create_group(req: GroupNameRequest): + try: + groups, group = watchlist.create_group(req.name, req.color) + except ValueError as e: + raise HTTPException(400, str(e)) from e + return {"groups": groups, "group": group} + + +@router.put("/groups/{group_id}") +def rename_group(group_id: str, req: GroupNameRequest): + try: + groups = watchlist.rename_group(group_id, req.name, req.color) + except KeyError as e: + raise HTTPException(404, "自选分组不存在") from e + except ValueError as e: + raise HTTPException(400, str(e)) from e + return {"groups": groups} + + +@router.delete("/groups/{group_id}") +def delete_group(group_id: str, request: Request): + try: + groups, rows = watchlist.delete_group(group_id) + except KeyError as e: + raise HTTPException(404, "自选分组不存在") from e + return {"groups": groups, "symbols": _with_names(rows, request)} @router.get("/ocr-status") @@ -131,6 +176,17 @@ def move_one_to_top(symbol: str, request: Request): return {"symbols": _with_names(rows, request)} +@router.put("/{symbol}/group") +def assign_group(symbol: str, req: GroupAssignRequest, request: Request): + try: + rows = watchlist.set_group(symbol, req.group_id) + except KeyError as e: + raise HTTPException(404, "自选标的不存在") from e + except ValueError as e: + raise HTTPException(400, str(e)) from e + return {"symbols": _with_names(rows, request)} + + @router.delete("/{symbol}") def remove_one(symbol: str, request: Request): rows = watchlist.remove(symbol) diff --git a/backend/app/services/watchlist.py b/backend/app/services/watchlist.py index 476a56c..e00bcdb 100644 --- a/backend/app/services/watchlist.py +++ b/backend/app/services/watchlist.py @@ -1,10 +1,17 @@ -"""自选股服务(§6.1)。 +"""自选股与分组服务。 -存储:`data/user_data/watchlist.parquet`,字段 symbol + added_at + note。 +自选存储于 ``data/user_data/watchlist.parquet``,分组定义存储于同目录的 +``watchlist_groups.json``。历史 Parquet 缺少 ``group_id`` 时按未分组读取。 """ from __future__ import annotations +import json import logging +import os +import threading +import uuid +from concurrent.futures import ThreadPoolExecutor +from concurrent.futures import TimeoutError as FuturesTimeout from datetime import datetime from pathlib import Path @@ -17,6 +24,30 @@ from app.tickflow.rate_limits import chunked, resolve_limit logger = logging.getLogger(__name__) +_LOCK = threading.RLock() +_MAX_GROUP_NAME_LENGTH = 24 +DEFAULT_GROUP_COLOR = "sky" +GROUP_COLORS = frozenset({ + "sky", + "blue", + "indigo", + "violet", + "fuchsia", + "rose", + "orange", + "amber", + "lime", + "emerald", + "teal", + "cyan", +}) +_ENTRY_SCHEMA = { + "symbol": pl.Utf8, + "added_at": pl.Utf8, + "note": pl.Utf8, + "group_id": pl.Utf8, +} + def _path() -> Path: p = settings.data_dir / "user_data" / "watchlist.parquet" @@ -24,70 +55,230 @@ def _path() -> Path: return p -def list_symbols() -> list[dict]: +def _groups_path() -> Path: + p = settings.data_dir / "user_data" / "watchlist_groups.json" + p.parent.mkdir(parents=True, exist_ok=True) + return p + + +def _empty_entries() -> pl.DataFrame: + return pl.DataFrame(schema=_ENTRY_SCHEMA) + + +def _read_entries() -> pl.DataFrame: p = _path() if not p.exists(): - return [] + return _empty_entries() df = pl.read_parquet(p) - if df.is_empty(): - return [] - return df.to_dicts() + defaults = {"symbol": "", "added_at": "", "note": "", "group_id": None} + for column, dtype in _ENTRY_SCHEMA.items(): + if column not in df.columns: + df = df.with_columns(pl.lit(defaults[column], dtype=dtype).alias(column)) + return df.select(list(_ENTRY_SCHEMA)) -def add(symbol: str, note: str = "") -> list[dict]: +def _write_entries(df: pl.DataFrame) -> None: p = _path() - if p.exists(): - df = pl.read_parquet(p) - # 已存在则先移除,后面重新插入到最前面 - if symbol in df["symbol"].to_list(): - df = df.filter(pl.col("symbol") != symbol) - else: - df = pl.DataFrame(schema={"symbol": pl.Utf8, "added_at": pl.Utf8, "note": pl.Utf8}) + tmp = p.with_suffix(p.suffix + ".tmp") + df.select(list(_ENTRY_SCHEMA)).write_parquet(tmp) + os.replace(tmp, p) - new_row = pl.DataFrame({ - "symbol": [symbol], - "added_at": [datetime.utcnow().isoformat(timespec="seconds")], - "note": [note], - }) - out = pl.concat([new_row, df], how="diagonal_relaxed") - out.write_parquet(p) - return out.to_dicts() + +def _read_groups() -> list[dict]: + p = _groups_path() + if not p.exists(): + return [] + try: + raw = json.loads(p.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError) as exc: + raise ValueError("自选分组配置损坏,请检查 watchlist_groups.json") from exc + if not isinstance(raw, list): + raise ValueError("自选分组配置格式不正确") + groups = [] + for item in raw: + if not isinstance(item, dict) or not item.get("id") or not item.get("name"): + continue + color = str(item.get("color", DEFAULT_GROUP_COLOR)) + groups.append({ + "id": str(item["id"]), + "name": str(item["name"]), + "color": color if color in GROUP_COLORS else DEFAULT_GROUP_COLOR, + }) + return groups + + +def _write_groups(groups: list[dict]) -> None: + p = _groups_path() + tmp = p.with_suffix(p.suffix + ".tmp") + tmp.write_text(json.dumps(groups, ensure_ascii=False, indent=2), encoding="utf-8") + os.replace(tmp, p) + + +def _normalize_group_name(name: str) -> str: + normalized = name.strip() + if not normalized: + raise ValueError("分组名称不能为空") + if len(normalized) > _MAX_GROUP_NAME_LENGTH: + raise ValueError(f"分组名称不能超过 {_MAX_GROUP_NAME_LENGTH} 个字符") + return normalized + + +def _normalize_group_color(color: str | None) -> str: + normalized = (color or DEFAULT_GROUP_COLOR).strip().lower() + if normalized not in GROUP_COLORS: + raise ValueError("不支持的分组颜色") + return normalized + + +def _validate_group_id(group_id: str | None, groups: list[dict]) -> None: + if group_id is not None and not any(group["id"] == group_id for group in groups): + raise ValueError("自选分组不存在") + + +def list_symbols() -> list[dict]: + with _LOCK: + df = _read_entries() + return [] if df.is_empty() else df.to_dicts() + + +def add(symbol: str, note: str = "", group_id: str | None = None) -> list[dict]: + rows, _ = add_batch([symbol], note=note, group_id=group_id) + return rows + + +def add_batch( + symbols: list[str], + note: str = "", + group_id: str | None = None, +) -> tuple[list[dict], int]: + """批量添加并保持既有语义:每个新处理的标的移动到列表最前面。""" + with _LOCK: + groups = _read_groups() + _validate_group_id(group_id, groups) + rows = _read_entries().to_dicts() + added = 0 + for symbol in symbols: + existing = next((row for row in rows if row["symbol"] == symbol), None) + if existing is None: + added += 1 + rows = [row for row in rows if row["symbol"] != symbol] + resolved_group_id = ( + group_id if group_id is not None else (existing or {}).get("group_id") + ) + rows.insert(0, { + "symbol": symbol, + "added_at": datetime.utcnow().isoformat(timespec="seconds"), + "note": note, + "group_id": resolved_group_id, + }) + out = pl.DataFrame(rows, schema=_ENTRY_SCHEMA) if rows else _empty_entries() + _write_entries(out) + return out.to_dicts(), added def remove(symbol: str) -> list[dict]: - p = _path() - if not p.exists(): - return [] - df = pl.read_parquet(p) - df = df.filter(pl.col("symbol") != symbol) - df.write_parquet(p) - return df.to_dicts() + with _LOCK: + df = _read_entries().filter(pl.col("symbol") != symbol) + _write_entries(df) + return df.to_dicts() def move_to_top(symbol: str) -> list[dict]: - p = _path() - if not p.exists(): - return [] - df = pl.read_parquet(p) - if df.is_empty() or symbol not in df["symbol"].to_list(): - return df.to_dicts() - target = df.filter(pl.col("symbol") == symbol) - rest = df.filter(pl.col("symbol") != symbol) - out = pl.concat([target, rest], how="diagonal_relaxed") - out.write_parquet(p) - return out.to_dicts() + with _LOCK: + df = _read_entries() + if df.is_empty() or symbol not in df["symbol"].to_list(): + return df.to_dicts() + target = df.filter(pl.col("symbol") == symbol) + rest = df.filter(pl.col("symbol") != symbol) + out = pl.concat([target, rest], how="diagonal_relaxed") + _write_entries(out) + return out.to_dicts() def clear() -> int: """清空自选列表。返回移除的数量。""" - p = _path() - if not p.exists(): - return 0 - df = pl.read_parquet(p) - count = df.height - if count > 0: - pl.DataFrame(schema={"symbol": pl.Utf8, "added_at": pl.Utf8, "note": pl.Utf8}).write_parquet(p) - return count + with _LOCK: + df = _read_entries() + count = df.height + if count > 0: + _write_entries(_empty_entries()) + return count + + +def list_groups() -> list[dict]: + with _LOCK: + return _read_groups() + + +def create_group(name: str, color: str | None = None) -> tuple[list[dict], dict]: + with _LOCK: + normalized = _normalize_group_name(name) + normalized_color = _normalize_group_color(color) + groups = _read_groups() + if any(group["name"].casefold() == normalized.casefold() for group in groups): + raise ValueError("分组名称已存在") + group = { + "id": uuid.uuid4().hex, + "name": normalized, + "color": normalized_color, + } + groups.append(group) + _write_groups(groups) + return groups, group + + +def rename_group(group_id: str, name: str, color: str | None = None) -> list[dict]: + with _LOCK: + normalized = _normalize_group_name(name) + groups = _read_groups() + target = next((group for group in groups if group["id"] == group_id), None) + if target is None: + raise KeyError(group_id) + if any( + group["id"] != group_id and group["name"].casefold() == normalized.casefold() + for group in groups + ): + raise ValueError("分组名称已存在") + target["name"] = normalized + if color is not None: + target["color"] = _normalize_group_color(color) + _write_groups(groups) + return groups + + +def delete_group(group_id: str) -> tuple[list[dict], list[dict]]: + """删除分组定义,原分组内的自选保留并转为未分组。""" + with _LOCK: + groups = _read_groups() + if not any(group["id"] == group_id for group in groups): + raise KeyError(group_id) + df = _read_entries().with_columns( + pl.when(pl.col("group_id") == group_id) + .then(None) + .otherwise(pl.col("group_id")) + .alias("group_id") + ) + remaining = [group for group in groups if group["id"] != group_id] + _write_entries(df) + _write_groups(remaining) + return remaining, df.to_dicts() + + +def set_group(symbol: str, group_id: str | None) -> list[dict]: + with _LOCK: + groups = _read_groups() + _validate_group_id(group_id, groups) + df = _read_entries() + if symbol not in df["symbol"].to_list(): + raise KeyError(symbol) + df = df.with_columns( + pl.when(pl.col("symbol") == symbol) + .then(pl.lit(group_id, dtype=pl.Utf8)) + .otherwise(pl.col("group_id")) + .alias("group_id") + ) + _write_entries(df) + return df.to_dicts() def fetch_quotes(symbols: list[str], capset: CapabilitySet, timeout_s: float = 8.0) -> list[dict]: @@ -96,8 +287,6 @@ def fetch_quotes(symbols: list[str], capset: CapabilitySet, timeout_s: float = 8 优先用 quote.batch;否则降级为 quote.by_symbol 单股请求。 timeout_s: 单批次请求超时(秒),防止 API 卡死阻塞整个请求。 """ - from concurrent.futures import ThreadPoolExecutor, TimeoutError as FuturesTimeout - if not symbols: return [] diff --git a/backend/tests/test_minute_range_api.py b/backend/tests/test_minute_range_api.py new file mode 100644 index 0000000..9f46a4b --- /dev/null +++ b/backend/tests/test_minute_range_api.py @@ -0,0 +1,107 @@ +"""多日分时 API 契约。""" + +import asyncio +from datetime import date, datetime +from types import SimpleNamespace +from unittest.mock import MagicMock + +import polars as pl +import pytest +from fastapi import HTTPException + +from app.api import kline as kline_api + + +def _request(repo=None, capset=None): + return SimpleNamespace( + app=SimpleNamespace( + state=SimpleNamespace( + repo=repo or MagicMock(), + capabilities=capset or MagicMock(), + ) + ) + ) + + +def test_minute_range_returns_latest_sessions_with_previous_closes(): + repo = MagicMock() + repo.resolve_asset_type.return_value = "stock" + repo.execute_one.return_value = ("浦发银行", 1.0, 1.0) + repo.get_minute_range.return_value = pl.DataFrame({ + "symbol": ["600000.SH"] * 3, + "datetime": [ + datetime(2026, 8, 5, 1, 30), + datetime(2026, 8, 6, 1, 30), + datetime(2026, 8, 7, 1, 30), + ], + "open": [10.0, 11.0, 12.0], + "high": [10.2, 11.2, 12.2], + "low": [9.8, 10.8, 11.8], + "close": [10.1, 11.1, 12.1], + "volume": [100.0, 110.0, 120.0], + "amount": [101_000.0, 122_100.0, 145_200.0], + }) + repo.get_daily_asset.return_value = pl.DataFrame({ + "date": [ + date(2026, 8, 4), + date(2026, 8, 5), + date(2026, 8, 6), + date(2026, 8, 7), + ], + "close": [9.9, 10.1, 11.1, 12.1], + }) + + result = kline_api.get_minute_range(_request(repo), "600000.SH", 2) + + assert result["name"] == "浦发银行" + assert result["requested_days"] == 2 + assert result["source"] == "local" + assert [session["date"] for session in result["sessions"]] == [ + "2026-08-06", + "2026-08-07", + ] + assert [session["prev_close"] for session in result["sessions"]] == [ + 10.1, + 11.1, + ] + assert result["sessions"][0]["rows"][0]["close"] == 11.1 + + +def test_minute_range_does_not_read_stock_store_for_index(): + repo = MagicMock() + repo.resolve_asset_type.return_value = "index" + repo.get_instruments_asset.return_value = pl.DataFrame() + + result = kline_api.get_minute_range(_request(repo), "000001.SH", 10) + + assert result["asset_type"] == "index" + assert result["sessions"] == [] + repo.get_minute_range.assert_not_called() + + +def test_sync_minute_single_uses_requested_days(monkeypatch): + repo = MagicMock() + repo.resolve_asset_type.return_value = "stock" + capset = MagicMock() + sync = MagicMock(return_value=2400) + refresh = MagicMock() + monkeypatch.setattr(kline_api, "_minute_allowed", lambda _: True) + monkeypatch.setattr(kline_api.kline_sync, "sync_and_persist_minute", sync) + monkeypatch.setattr("app.jobs.daily_pipeline._refresh_single_view", refresh) + + result = asyncio.run(kline_api.sync_minute_single( + _request(repo, capset), + {"symbol": "600000.SH", "days": 10}, + )) + + assert result["rows"] == 2400 + sync.assert_called_once_with(["600000.SH"], repo, capset, days=10) + refresh.assert_called_once_with(repo, "kline_minute") + + +def test_sync_minute_single_rejects_invalid_days(): + with pytest.raises(HTTPException, match="days 必须在 1 到 30 之间"): + asyncio.run(kline_api.sync_minute_single( + _request(), + {"symbol": "600000.SH", "days": 0}, + )) diff --git a/backend/tests/test_watchlist_groups.py b/backend/tests/test_watchlist_groups.py new file mode 100644 index 0000000..83b550e --- /dev/null +++ b/backend/tests/test_watchlist_groups.py @@ -0,0 +1,114 @@ +"""自选分组持久化与 API 契约。""" +from types import SimpleNamespace +from unittest.mock import MagicMock + +import polars as pl +import pytest +from fastapi import HTTPException + +from app.api import watchlist as watchlist_api +from app.config import settings +from app.services import watchlist + + +def _request(): + repo = MagicMock() + repo.get_name_map.return_value = {} + return SimpleNamespace(app=SimpleNamespace(state=SimpleNamespace(repo=repo))) + + +def test_historical_watchlist_is_read_as_ungrouped(monkeypatch, tmp_path): + monkeypatch.setattr(settings, "data_dir", tmp_path) + path = tmp_path / "user_data" / "watchlist.parquet" + path.parent.mkdir(parents=True) + pl.DataFrame({ + "symbol": ["600000.SH"], + "added_at": ["2026-08-08T10:00:00"], + "note": [""], + }).write_parquet(path) + + assert watchlist.list_symbols()[0]["group_id"] is None + + +def test_group_lifecycle_preserves_watchlist_entries(monkeypatch, tmp_path): + monkeypatch.setattr(settings, "data_dir", tmp_path) + groups, created = watchlist.create_group(" 短线 ", "orange") + assert groups == [{"id": created["id"], "name": "短线", "color": "orange"}] + + watchlist.add("600000.SH", group_id=created["id"]) + watchlist.add("000001.SZ") + assert watchlist.list_symbols()[1]["group_id"] == created["id"] + + renamed = watchlist.rename_group(created["id"], "观察", "fuchsia") + assert renamed[0]["name"] == "观察" + assert renamed[0]["color"] == "fuchsia" + + remaining, rows = watchlist.delete_group(created["id"]) + assert remaining == [] + assert {row["symbol"] for row in rows} == {"600000.SH", "000001.SZ"} + assert all(row["group_id"] is None for row in rows) + + +def test_group_validation_and_assignment_errors(monkeypatch, tmp_path): + monkeypatch.setattr(settings, "data_dir", tmp_path) + _, created = watchlist.create_group("核心") + watchlist.add("600000.SH") + + with pytest.raises(ValueError, match="已存在"): + watchlist.create_group("核心") + with pytest.raises(ValueError, match="颜色"): + watchlist.create_group("无效颜色", "black") + with pytest.raises(ValueError, match="颜色"): + watchlist.rename_group(created["id"], "核心", "black") + with pytest.raises(ValueError, match="不存在"): + watchlist.set_group("600000.SH", "missing") + with pytest.raises(KeyError): + watchlist.set_group("000001.SZ", created["id"]) + + rows = watchlist.set_group("600000.SH", created["id"]) + assert rows[0]["group_id"] == created["id"] + rows = watchlist.set_group("600000.SH", None) + assert rows[0]["group_id"] is None + + +def test_group_api_contract(monkeypatch, tmp_path): + monkeypatch.setattr(settings, "data_dir", tmp_path) + request = _request() + created = watchlist_api.create_group( + watchlist_api.GroupNameRequest(name="中线", color="teal") + ) + group_id = created["group"]["id"] + assert created["group"]["color"] == "teal" + + added = watchlist_api.add_one( + watchlist_api.AddRequest(symbol="600000.SH", group_id=group_id), + request, + ) + assert added["symbols"][0]["group_id"] == group_id + + moved = watchlist_api.assign_group( + "600000.SH", + watchlist_api.GroupAssignRequest(group_id=None), + request, + ) + assert moved["symbols"][0]["group_id"] is None + + with pytest.raises(HTTPException) as exc_info: + watchlist_api.rename_group("missing", watchlist_api.GroupNameRequest(name="无效")) + assert exc_info.value.status_code == 404 + + +def test_historical_groups_default_to_sky(monkeypatch, tmp_path): + monkeypatch.setattr(settings, "data_dir", tmp_path) + path = tmp_path / "user_data" / "watchlist_groups.json" + path.parent.mkdir(parents=True) + path.write_text( + '[{"id":"legacy","name":"旧分组"},' + '{"id":"invalid","name":"未知颜色","color":"black"}]', + encoding="utf-8", + ) + + assert watchlist.list_groups() == [ + {"id": "legacy", "name": "旧分组", "color": "sky"}, + {"id": "invalid", "name": "未知颜色", "color": "sky"}, + ] diff --git a/backend/uv.lock b/backend/uv.lock index f38c899..6b6151f 100644 --- a/backend/uv.lock +++ b/backend/uv.lock @@ -1958,6 +1958,15 @@ wheels = [ { url = "https://pypi.tuna.tsinghua.edu.cn/packages/10/bd/c038d7cc38edc1aa5bf91ab8068b63d4308c66c4c8bb3cbba7dfbc049f9c/pyparsing-3.3.2-py3-none-any.whl", hash = "sha256:850ba148bd908d7e2411587e247a1e4f0327839c40e2e5e6d05a007ecc69911d", size = 122781, upload-time = "2026-01-21T03:57:55.912Z" }, ] +[[package]] +name = "pypinyin" +version = "0.55.0" +source = { registry = "https://pypi.tuna.tsinghua.edu.cn/simple" } +sdist = { url = "https://pypi.tuna.tsinghua.edu.cn/packages/b4/a4/784cf98c09e0dc22776b0d7d8a4a5b761218bcae4608c2416ce1e167c8af/pypinyin-0.55.0.tar.gz", hash = "sha256:b5711b3a0c6f76e67408ec6b2e3c4987a3a806b7c528076e7c7b86fcf0eaa66b", size = 839836, upload-time = "2025-07-20T12:01:50.657Z" } +wheels = [ + { url = "https://pypi.tuna.tsinghua.edu.cn/packages/b9/7b/4cabc76fcc21c3c7d5c671d8783984d30ac9d3bb387c4ba784fca3cdfa3a/pypinyin-0.55.0-py2.py3-none-any.whl", hash = "sha256:d53b1e8ad2cdb815fb2cb604ed3123372f5a28c6f447571244aca36fc62a286f", size = 840203, upload-time = "2025-07-20T12:01:48.535Z" }, +] + [[package]] name = "pytesseract" version = "0.3.13" @@ -2504,7 +2513,7 @@ all = [ [[package]] name = "tickflow-stock-panel-backend" -version = "0.1.83" +version = "0.1.88" source = { editable = "." } dependencies = [ { name = "apscheduler" }, @@ -2523,6 +2532,7 @@ dependencies = [ { name = "pyarrow" }, { name = "pydantic" }, { name = "pydantic-settings" }, + { name = "pypinyin" }, { name = "pytesseract" }, { name = "python-dotenv" }, { name = "python-multipart" }, @@ -2570,6 +2580,7 @@ requires-dist = [ { name = "pyarrow", specifier = ">=16.0" }, { name = "pydantic", specifier = ">=2.7" }, { name = "pydantic-settings", specifier = ">=2.4" }, + { name = "pypinyin", specifier = ">=0.50" }, { name = "pytesseract", specifier = ">=0.3.10" }, { name = "pytest", marker = "extra == 'dev'", specifier = ">=8.0" }, { name = "pytest-asyncio", marker = "extra == 'dev'", specifier = ">=0.23" }, diff --git a/frontend/src/components/EChartsIntraday.tsx b/frontend/src/components/EChartsIntraday.tsx index b4ba6a8..c308f01 100644 --- a/frontend/src/components/EChartsIntraday.tsx +++ b/frontend/src/components/EChartsIntraday.tsx @@ -2,6 +2,7 @@ import { useEffect, useMemo, useRef, useState } from 'react' import * as echarts from 'echarts' import type { ECharts, EChartsOption } from 'echarts' import type { MinuteKlineRow, PriceLimitInfo } from '@/lib/api' +import { computeIntradayAverage, formatMinuteTime, FULL_DAY_TIMES } from '@/lib/intraday-chart' import { useChartTheme, type ChartTheme } from '@/lib/theme' type YMode = 'adaptive' | 'limit' @@ -26,26 +27,6 @@ interface Props { showAvgLine?: boolean } -function fmtTime(dt: string): string { - const match = dt.match(/(\d{2}):(\d{2})/) - if (!match) return dt.slice(11, 16) - const h = (parseInt(match[1]) + 8) % 24 - return `${String(h).padStart(2, '0')}:${match[2]}` -} - -function computeAvgPrice(data: MinuteKlineRow[]): number[] { - // 分时均线 = 累计成交额 / 累计成交量(手→股) - const result: number[] = [] - let sumAmt = 0 - let sumVol = 0 - for (const d of data) { - sumAmt += d.amount - sumVol += d.volume * 100 - result.push(sumVol > 0 ? sumAmt / sumVol : d.close) - } - return result -} - function fmtAmt(v: number): string { if (v >= 1_000_000_000) return `${(v / 1_000_000_000).toFixed(2)}亿` if (v >= 10_000) return `${(v / 10_000).toFixed(0)}万` @@ -56,29 +37,6 @@ function isValidPrice(v: number | null | undefined): v is number { return typeof v === 'number' && Number.isFinite(v) && v > 0 } -/** 生成全天分时时间刻度 9:30 ~ 11:30, 13:00 ~ 15:00, 每分钟一个点 (共242个) */ -function generateFullDayTimes(): string[] { - const times: string[] = [] - // 上午 9:30 ~ 11:30 (121 分钟) - for (let h = 9; h <= 11; h++) { - const startM = h === 9 ? 30 : 0 - const endM = h === 11 ? 30 : 59 - for (let m = startM; m <= endM; m++) { - times.push(`${String(h).padStart(2, '0')}:${String(m).padStart(2, '0')}`) - } - } - // 下午 13:00 ~ 15:00 (121 分钟) - for (let h = 13; h <= 15; h++) { - const endM = h === 15 ? 0 : 59 - for (let m = 0; m <= endM; m++) { - times.push(`${String(h).padStart(2, '0')}:${String(m).padStart(2, '0')}`) - } - } - return times -} - -const FULL_DAY_TIMES = generateFullDayTimes() - /** 计算实际涨跌停价 (四舍五入到2位小数) 和实际涨跌停幅度 */ function getLimitPrices(prevClose: number, priceLimit?: PriceLimitInfo): { limitUp: number // 涨停价 (四舍五入) @@ -112,7 +70,7 @@ function buildOption(data: MinuteKlineRow[], prevClose: number | undefined, avgP const volNeutral = 'rgba(161,161,170,0.5)' for (let i = 0; i < data.length; i++) { - const timeKey = fmtTime(data[i].datetime) + const timeKey = formatMinuteTime(data[i].datetime) const idx = timeIndexMap.get(timeKey) if (idx !== undefined) { closes[idx] = data[i].close @@ -414,7 +372,7 @@ export function EChartsIntraday({ data, height = 320, prevClose, date, priceLimi const [infoIdx, setInfoIdx] = useState(data.length - 1) const [yMode, setYMode] = useState('adaptive') const ct = useChartTheme() - const avgPrices = useMemo(() => computeAvgPrice(data), [data]) + const avgPrices = useMemo(() => computeIntradayAverage(data), [data]) // 分时线颜色:基于最新价 vs 昨收 const lastClose = data.length > 0 ? data[data.length - 1].close : null @@ -480,7 +438,7 @@ export function EChartsIntraday({ data, height = 320, prevClose, date, priceLimi const timeIndexMap = new Map(FULL_DAY_TIMES.map((t, i) => [t, i])) const mapping = new Map() for (let i = 0; i < data.length; i++) { - const timeKey = fmtTime(data[i].datetime) + const timeKey = formatMinuteTime(data[i].datetime) const fullDayIdx = timeIndexMap.get(timeKey) if (fullDayIdx !== undefined) { mapping.set(fullDayIdx, i) diff --git a/frontend/src/components/EChartsMultiDayIntraday.tsx b/frontend/src/components/EChartsMultiDayIntraday.tsx new file mode 100644 index 0000000..32f8a22 --- /dev/null +++ b/frontend/src/components/EChartsMultiDayIntraday.tsx @@ -0,0 +1,373 @@ +import { useEffect, useMemo, useRef, useState } from 'react' +import * as echarts from 'echarts' +import type { ECharts, EChartsOption } from 'echarts' +import type { MinuteKlineRow, MinuteKlineSession } from '@/lib/api' +import { computeIntradayAverage, formatMinuteTime, FULL_DAY_TIMES } from '@/lib/intraday-chart' +import { useChartTheme } from '@/lib/theme' + +const COLORS = { + up: '#C74040', + down: '#2D9B65', + flat: '#A1A1AA', + average: '#F59E0B', + volumeUp: 'rgba(240,68,56,0.58)', + volumeDown: 'rgba(18,183,106,0.58)', + volumeFlat: 'rgba(161,161,170,0.45)', +} + +interface Props { + sessions: MinuteKlineSession[] + height?: number +} + +interface InfoPoint { + date: string + row: MinuteKlineRow + average: number + prevClose: number | null +} + +function formatAmount(value: number): string { + if (value >= 1_000_000_000) return `${(value / 1_000_000_000).toFixed(2)}亿` + if (value >= 10_000) return `${(value / 10_000).toFixed(0)}万` + return value.toFixed(0) +} + +function priceColor(close: number, prevClose: number | null): string { + if (prevClose == null || close === prevClose) return COLORS.flat + return close > prevClose ? COLORS.up : COLORS.down +} + +function buildModel(sessions: MinuteKlineSession[]) { + const categories: string[] = [] + const volumeData: ({ value: number; itemStyle: { color: string } } | null)[] = [] + const dayLabelByIndex = new Map() + const dayStartIndexes: number[] = [] + const pointByIndex = new Map() + const dayRanges: { + start: number + session: MinuteKlineSession + values: (number | null)[] + averages: (number | null)[] + }[] = [] + const priceValues: number[] = [] + + const labelStep = Math.max(1, Math.ceil(sessions.length / 10)) + for (let sessionIndex = 0; sessionIndex < sessions.length; sessionIndex++) { + const session = sessions[sessionIndex] + const start = categories.length + dayStartIndexes.push(start) + if (sessionIndex % labelStep === 0 || sessionIndex === sessions.length - 1) { + dayLabelByIndex.set(start + Math.floor(FULL_DAY_TIMES.length / 2), session.date.slice(5)) + } + + const averagePrices = computeIntradayAverage(session.rows) + const rowsByTime = new Map() + session.rows.forEach((row, index) => { + rowsByTime.set(formatMinuteTime(row.datetime), { + row, + average: averagePrices[index], + }) + }) + + const dayValues: (number | null)[] = [] + const dayAverages: (number | null)[] = [] + for (const time of FULL_DAY_TIMES) { + const point = rowsByTime.get(time) + const index = categories.length + categories.push(`${session.date} ${time}`) + if (!point) { + dayValues.push(null) + dayAverages.push(null) + volumeData.push(null) + continue + } + + const { row, average } = point + dayValues.push(row.close) + dayAverages.push(average) + volumeData.push({ + value: row.volume, + itemStyle: { + color: row.close > row.open + ? COLORS.volumeUp + : row.close < row.open + ? COLORS.volumeDown + : COLORS.volumeFlat, + }, + }) + priceValues.push(row.low, row.high, average) + pointByIndex.set(index, { + date: session.date, + row, + average, + prevClose: session.prev_close, + }) + } + + dayRanges.push({ + start, + session, + values: dayValues, + averages: dayAverages, + }) + + if (sessionIndex < sessions.length - 1) { + categories.push(`${session.date} gap`) + volumeData.push(null) + } + } + + return { + categories, + volumeData, + dayLabelByIndex, + dayStartIndexes, + pointByIndex, + dayRanges, + priceValues, + latest: pointByIndex.size > 0 + ? Array.from(pointByIndex.values())[pointByIndex.size - 1] + : null, + } +} + +export function EChartsMultiDayIntraday({ sessions, height = 420 }: Props) { + const containerRef = useRef(null) + const chartRef = useRef(null) + const resizeObserverRef = useRef(null) + const model = useMemo(() => buildModel(sessions), [sessions]) + const modelRef = useRef(model) + modelRef.current = model + const [info, setInfo] = useState(model.latest) + const theme = useChartTheme() + + useEffect(() => { + setInfo(model.latest) + }, [model]) + + useEffect(() => { + const container = containerRef.current + if (!container) return + + let chart = chartRef.current + if (!chart) { + chart = echarts.init(container, undefined, { renderer: 'canvas' }) + chartRef.current = chart + resizeObserverRef.current = new ResizeObserver(() => chart?.resize()) + resizeObserverRef.current.observe(container) + + chart.on('updateAxisPointer', (event: any) => { + const axisInfo = event.axesInfo?.find((item: any) => item.axisDim === 'x' && item.axisIndex === 0) + ?? event.axesInfo?.[0] + const rawValue = axisInfo?.value + const current = modelRef.current + const index = typeof rawValue === 'number' + ? rawValue + : current.categories.indexOf(String(rawValue)) + const point = current.pointByIndex.get(index) + if (point) setInfo(point) + }) + chart.on('globalout', () => setInfo(modelRef.current.latest)) + } + + const minPrice = model.priceValues.length > 0 ? Math.min(...model.priceValues) : 0 + const maxPrice = model.priceValues.length > 0 ? Math.max(...model.priceValues) : 1 + const padding = Math.max((maxPrice - minPrice) * 0.08, maxPrice * 0.002) + const totalLength = model.categories.length + const priceSeries: any[] = model.dayRanges.map(({ start, session, values }) => { + const data = new Array(totalLength).fill(null) as (number | null)[] + for (let index = 0; index < values.length; index++) data[start + index] = values[index] + const last = session.rows[session.rows.length - 1] + const color = last ? priceColor(last.close, session.prev_close) : COLORS.flat + return { + name: session.date, + type: 'line', + data, + symbol: 'none', + smooth: false, + connectNulls: true, + lineStyle: { width: 1.2, color }, + areaStyle: { color, opacity: 0.08 }, + emphasis: { disabled: true }, + } + }) + + const boundaryData = model.dayStartIndexes.slice(1).map(index => ({ + xAxis: model.categories[index], + lineStyle: { color: theme.grid, width: 1 }, + label: { show: false }, + })) + if (priceSeries.length > 0 && boundaryData.length > 0) { + priceSeries[0].markLine = { + symbol: 'none', + silent: true, + data: boundaryData, + } + } + const averageSeries: any[] = model.dayRanges.map(({ start, session, averages }) => { + const data = new Array(totalLength).fill(null) as (number | null)[] + for (let index = 0; index < averages.length; index++) data[start + index] = averages[index] + return { + name: `${session.date} 均价`, + type: 'line', + data, + symbol: 'none', + connectNulls: true, + lineStyle: { width: 1, color: COLORS.average }, + emphasis: { disabled: true }, + } + }) + const option: EChartsOption = { + animation: false, + backgroundColor: 'transparent', + tooltip: { + trigger: 'axis', + backgroundColor: 'transparent', + borderWidth: 0, + formatter: () => '', + axisPointer: { + type: 'cross', + label: { + show: true, + backgroundColor: theme.tooltipBg, + borderColor: theme.tooltipBorder, + borderWidth: 1, + color: theme.tooltipText, + fontFamily: 'JetBrains Mono, monospace', + fontSize: 10, + }, + crossStyle: { color: theme.crosshair, type: 'dashed', width: 1 }, + lineStyle: { color: theme.crosshair, type: 'dashed', width: 1 }, + }, + }, + axisPointer: { link: [{ xAxisIndex: 'all' }] }, + grid: [ + { left: 58, right: 18, top: 16, bottom: '28%' }, + { left: 58, right: 18, top: '76%', bottom: 22 }, + ], + xAxis: [ + { + type: 'category', + data: model.categories, + boundaryGap: false, + axisLine: { lineStyle: { color: theme.grid } }, + axisTick: { show: false }, + splitLine: { show: false }, + axisLabel: { + color: theme.text, + fontFamily: 'JetBrains Mono, monospace', + fontSize: 10, + interval: 0, + hideOverlap: true, + formatter: (_value: string, index: number) => model.dayLabelByIndex.get(index) ?? '', + }, + axisPointer: { + label: { + formatter: (params: any) => { + const value = String(params.value ?? '') + return value.endsWith(' gap') ? '' : value.slice(5) + }, + }, + }, + }, + { + type: 'category', + gridIndex: 1, + data: model.categories, + boundaryGap: false, + axisLine: { show: false }, + axisTick: { show: false }, + axisLabel: { show: false }, + splitLine: { show: false }, + }, + ], + yAxis: [ + { + type: 'value', + min: minPrice - padding, + max: maxPrice + padding, + scale: true, + axisLine: { show: false }, + axisTick: { show: false }, + splitLine: { lineStyle: { color: theme.grid } }, + axisLabel: { + color: theme.text, + fontFamily: 'JetBrains Mono, monospace', + fontSize: 10, + formatter: (value: number) => value.toFixed(2), + }, + }, + { + type: 'value', + gridIndex: 1, + scale: true, + axisLine: { show: false }, + axisTick: { show: false }, + splitLine: { show: false }, + axisLabel: { show: false }, + }, + ], + dataZoom: [{ + type: 'inside', + xAxisIndex: [0, 1], + start: 0, + end: 100, + minValueSpan: FULL_DAY_TIMES.length, + filterMode: 'none', + }], + series: [ + ...priceSeries, + ...averageSeries, + { + name: '成交量', + type: 'bar', + data: model.volumeData, + xAxisIndex: 1, + yAxisIndex: 1, + }, + ], + } + chart.setOption(option, true) + }, [height, model, theme]) + + useEffect(() => () => { + chartRef.current?.off('updateAxisPointer') + chartRef.current?.off('globalout') + resizeObserverRef.current?.disconnect() + chartRef.current?.dispose() + chartRef.current = null + }, []) + + const changePct = info?.prevClose + ? (info.row.close - info.prevClose) / info.prevClose * 100 + : null + const infoColor = info ? priceColor(info.row.close, info.prevClose) : COLORS.flat + const rowCount = sessions.reduce((total, session) => total + session.rows.length, 0) + + return ( +
+
+
+ {info ? ( + <> + {info.date} {formatMinuteTime(info.row.datetime)} + {info.row.open.toFixed(2)} + {info.row.high.toFixed(2)} + {info.row.low.toFixed(2)} + {info.row.close.toFixed(2)} + {changePct != null && ( + {changePct >= 0 ? '+' : ''}{changePct.toFixed(2)}% + )} + 均价{info.average.toFixed(2)} + {info.row.volume.toFixed(0)} + {formatAmount(info.row.amount)} + + ) : } +
+
{sessions.length} 个交易日 · {rowCount} 分钟
+
+
+
+ ) +} diff --git a/frontend/src/components/StockInfoBar.tsx b/frontend/src/components/StockInfoBar.tsx index 1063249..7b5891f 100644 --- a/frontend/src/components/StockInfoBar.tsx +++ b/frontend/src/components/StockInfoBar.tsx @@ -3,6 +3,7 @@ import { Settings2, RadioTower, Star } from 'lucide-react' import type { KlineRow, FinancialMetricRecord } from '@/lib/api' import { fmtPrice, fmtBigNum, fmtVolume } from '@/lib/format' import { ListColumnCustomizer } from '@/components/ListColumnCustomizer' +import { WatchlistAddMenu } from '@/components/WatchlistAddMenu' import { INFO_GROUPS, type ColumnConfig } from '@/lib/stock-info-fields' const BULL = '#C74040' @@ -20,9 +21,11 @@ interface Props { financialMetrics?: FinancialMetricRecord /** 加监控回调 (个股弹窗传入, 有值时渲染 RadioTower 图标) */ onMonitor?: () => void - /** 加自选回调 + 是否已自选 (有 onToggle 时渲染 Star 图标) */ + /** 自选状态与操作(传入对应回调时渲染 Star 图标) */ inWatchlist?: boolean - onToggleWatchlist?: () => void + onAddToWatchlist?: (groupId: string | null) => void + onRemoveFromWatchlist?: () => void + watchlistPending?: boolean } /** @@ -91,7 +94,20 @@ function renderExtInline( ) } -export function StockInfoBar({ symbol, name, stockInfo, rows, fields, onFieldsChange, financialMetrics, onMonitor, inWatchlist, onToggleWatchlist }: Props) { +export function StockInfoBar({ + symbol, + name, + stockInfo, + rows, + fields, + onFieldsChange, + financialMetrics, + onMonitor, + inWatchlist, + onAddToWatchlist, + onRemoveFromWatchlist, + watchlistPending, +}: Props) { // 弹窗开关:纯本地状态,与数据/配置无关,放早期 return 之前 const [customizerOpen, setCustomizerOpen] = useState(false) // ext 标签展开状态:按 symbol::colId,切股/切字段时互不干扰 @@ -216,15 +232,27 @@ export function StockInfoBar({ symbol, name, stockInfo, rows, fields, onFieldsCh {/* 右侧操作按钮:加自选 + 加监控 + 信息条配置 */}
- {onToggleWatchlist && ( + {inWatchlist && onRemoveFromWatchlist ? ( - )} + ) : !inWatchlist && onAddToWatchlist ? ( + + + + ) : null} {onMonitor && ( +
+ ) + } + + if (sessions.length === 0) { + return ( +
+ {syncMinute.isPending ? ( + <> + + 正在获取近 {days} 日分钟 K… + + ) : ( + <> + {isIndex ? '指数暂无分钟数据' : '本地暂无可展示的分钟数据'} + {!isIndex && ( + + )} + + )} + {syncMinute.isError && {errorMessage(syncMinute.error)}} +
+ ) + } + + return ( +
+ {showCoverage && ( +
+ 当前有 {sessions.length} 个交易日数据,目标 {days} 日 + +
+ )} + + {syncMinute.isError && ( +
{errorMessage(syncMinute.error)}
+ )} +
+ ) +} diff --git a/frontend/src/components/StockPanel.tsx b/frontend/src/components/StockPanel.tsx index e6e496c..4e4f707 100644 --- a/frontend/src/components/StockPanel.tsx +++ b/frontend/src/components/StockPanel.tsx @@ -29,9 +29,11 @@ interface Props { showMarkerToggle?: boolean /** 加监控回调 (传入后信息条显示 RadioTower 图标) */ onMonitor?: () => void - /** 加自选 (传入后信息条显示 Star 图标) */ + /** 自选操作(传入后信息条显示 Star 图标) */ inWatchlist?: boolean - onToggleWatchlist?: () => void + onAddToWatchlist?: (groupId: string | null) => void + onRemoveFromWatchlist?: () => void + watchlistPending?: boolean /** 分时图自动刷新间隔(ms)。undefined = 不轮询。个股对话框盘中实时刷新时传入。 */ refetchIntervalMs?: number } @@ -52,7 +54,9 @@ export function StockPanel({ showMarkerToggle = true, onMonitor, inWatchlist, - onToggleWatchlist, + onAddToWatchlist, + onRemoveFromWatchlist, + watchlistPending, refetchIntervalMs, }: Props) { const [linkedPrice, setLinkedPrice] = useState(null) @@ -132,7 +136,9 @@ export function StockPanel({ financialMetrics={financialMetrics} onMonitor={onMonitor} inWatchlist={inWatchlist} - onToggleWatchlist={onToggleWatchlist} + onAddToWatchlist={onAddToWatchlist} + onRemoveFromWatchlist={onRemoveFromWatchlist} + watchlistPending={watchlistPending} />
diff --git a/frontend/src/components/StockPreviewDialog.tsx b/frontend/src/components/StockPreviewDialog.tsx index f49a3ab..874b189 100644 --- a/frontend/src/components/StockPreviewDialog.tsx +++ b/frontend/src/components/StockPreviewDialog.tsx @@ -1,16 +1,18 @@ import { useState, useEffect } from 'react' import { useQuery, useMutation, useQueryClient } from '@tanstack/react-query' import { motion, AnimatePresence } from 'framer-motion' -import { X, RefreshCw, Clock } from 'lucide-react' +import { X, RefreshCw, Clock, LineChart } from 'lucide-react' import { api } from '@/lib/api' import { QK } from '@/lib/queryKeys' import { cnSignal } from '@/lib/signals' import { StockPanel, getDefaultRange } from '@/components/StockPanel' +import { StockMultiDayIntradayChart } from '@/components/StockMultiDayIntradayChart' import { DatePicker } from '@/components/DatePicker' import { RuleEditor } from '@/components/monitor/RuleEditor' import { usePreferences, useQuoteStatus } from '@/lib/useSharedQueries' import { setFocusSymbol, clearFocusSymbol } from '@/lib/useQuoteStream' import { useDialogBackdrop } from '@/lib/useDialogBackdrop' +import { storage } from '@/lib/storage' interface Props { symbol: string | null @@ -34,6 +36,16 @@ const PRESETS: { label: string; months: number }[] = [ { label: '1年', months: 12 }, ] +type PreviewView = 'daily' | 'intraday' +const INTRADAY_DAY_OPTIONS = [1, 5, 10, 20] as const + +function loadIntradayDays(): number { + const saved = storage.stockPreviewIntradayDays.get(10) + return INTRADAY_DAY_OPTIONS.includes(saved as typeof INTRADAY_DAY_OPTIONS[number]) + ? saved + : 10 +} + function boardTag(symbol: string): { label: string; color: string } | null { if (/^(300|301)/.test(symbol)) return { label: '创', color: 'text-[#f97316] bg-[#f97316]/12 border-[#f97316]/25' } if (/^688/.test(symbol)) return { label: '科', color: 'text-purple-400 bg-purple-400/12 border-purple-400/25' } @@ -42,7 +54,8 @@ function boardTag(symbol: string): { label: string; color: string } | null { } export function StockPreviewDialog({ symbol, name, onClose, triggerInfo }: Props) { - const [showIntraday, setShowIntraday] = useState(false) + const [view, setView] = useState('daily') + const [intradayDays, setIntradayDays] = useState(loadIntradayDays) const [dateRange, setDateRange] = useState(getDefaultRange) const [showMonitorEditor, setShowMonitorEditor] = useState(false) const qc = useQueryClient() @@ -56,7 +69,15 @@ export function StockPreviewDialog({ symbol, name, onClose, triggerInfo }: Props const inWatchlist = (watchlist.data?.symbols ?? []).some((s: any) => s.symbol === symbol) const toggleWatchlist = useMutation({ - mutationFn: () => inWatchlist ? api.watchlistRemove(symbol!) : api.watchlistAdd(symbol!), + mutationFn: ({ + action, + groupId, + }: { + action: 'add' | 'remove' + groupId?: string | null + }) => action === 'remove' + ? api.watchlistRemove(symbol!) + : api.watchlistAdd(symbol!, '', groupId), onSuccess: () => { qc.invalidateQueries({ queryKey: QK.watchlist }) qc.invalidateQueries({ queryKey: ['watchlist-enriched'] }) @@ -73,6 +94,10 @@ export function StockPreviewDialog({ symbol, name, onClose, triggerInfo }: Props return () => document.removeEventListener('keydown', handler) }, [symbol, onClose]) + useEffect(() => { + if (symbol) setView('daily') + }, [symbol]) + // 焦点股票注册: SSE quotes_updated 推送时精准 invalidate 当前股票日K, // 让对话框日K最后一根蜡烛随实时价变化 (后端只读内存, 不调 TickFlow)。 // 关闭/切股时清除, 避免无谓刷新。 @@ -94,12 +119,19 @@ export function StockPreviewDialog({ symbol, name, onClose, triggerInfo }: Props const handleRefresh = () => { if (!symbol) return - qc.invalidateQueries({ queryKey: ['kline', symbol!] }) - if (showIntraday) { + if (view === 'daily') { + qc.invalidateQueries({ queryKey: ['kline', symbol] }) + } else { + qc.invalidateQueries({ queryKey: ['kline-minute-range', symbol] }) qc.invalidateQueries({ queryKey: ['kline-minute', symbol!] }) } } + const selectIntradayDays = (days: number) => { + setIntradayDays(days) + storage.stockPreviewIntradayDays.set(days) + } + return ( {symbol && ( @@ -123,8 +155,8 @@ export function StockPreviewDialog({ symbol, name, onClose, triggerInfo }: Props className="relative w-[92vw] max-w-[1100px] max-h-[95vh] rounded-card border border-border bg-base shadow-2xl overflow-hidden flex flex-col" > {/* 顶栏 */} -
-
+
+
{(() => { const board = symbol ? boardTag(symbol) : null return board ? ( @@ -133,11 +165,51 @@ export function StockPreviewDialog({ symbol, name, onClose, triggerInfo }: Props ) : null })()} - {symbol} - {name && {name}} + {symbol} + {name && {name}}
-
+ +
+ +
+
+ + +
+ +
+ {view === 'daily' ? ( + <> {/* 日期范围快捷 */} {PRESETS.map(p => { const now = new Date() @@ -175,23 +247,31 @@ export function StockPreviewDialog({ symbol, name, onClose, triggerInfo }: Props onChange={(v) => setDateRange(prev => ({ ...prev, end: v }))} min={dateRange.start} /> + + ) : ( + <> + 区间 +
+ {INTRADAY_DAY_OPTIONS.map(days => ( + + ))} +
+ + )} - | - - {/* 分时开关 */} - - - | + {/* 刷新 */} - - {/* 关闭 */} -
@@ -253,19 +325,28 @@ export function StockPreviewDialog({ symbol, name, onClose, triggerInfo }: Props
)} - {/* K 线内容 */} + {/* 图表内容 */}
- { if (!showIntraday) setShowIntraday(true) }} - dateRange={dateRange} - onMonitor={() => setShowMonitorEditor(true)} - inWatchlist={inWatchlist} - onToggleWatchlist={() => toggleWatchlist.mutate()} - refetchIntervalMs={intradayRefetchMs} - /> + {view === 'daily' ? ( + setShowMonitorEditor(true)} + inWatchlist={inWatchlist} + onAddToWatchlist={groupId => toggleWatchlist.mutate({ action: 'add', groupId })} + onRemoveFromWatchlist={() => toggleWatchlist.mutate({ action: 'remove' })} + watchlistPending={toggleWatchlist.isPending} + /> + ) : ( + + )}
{/* 加监控编辑器弹层 */} diff --git a/frontend/src/components/WatchlistAddMenu.tsx b/frontend/src/components/WatchlistAddMenu.tsx new file mode 100644 index 0000000..bea247d --- /dev/null +++ b/frontend/src/components/WatchlistAddMenu.tsx @@ -0,0 +1,268 @@ +import { useEffect, useId, useLayoutEffect, useRef, useState, type ReactNode } from 'react' +import { createPortal } from 'react-dom' +import { useQuery } from '@tanstack/react-query' +import { Check, Folder, Inbox, List, LoaderCircle, RefreshCw } from 'lucide-react' +import { api } from '@/lib/api' +import { QK } from '@/lib/queryKeys' +import { resolveWatchlistGroupColor } from '@/lib/watchlist-group-colors' + +const MENU_WIDTH = 224 +const MENU_MAX_HEIGHT = 320 +const VIEWPORT_GAP = 8 +const TRIGGER_GAP = 6 + +export interface WatchlistGroupMenuProps { + children: ReactNode + onSelect: (groupId: string | null) => void + disabled?: boolean + preferredGroupId?: string | null + includeAll?: boolean + counts?: Record + total?: number + disableEmpty?: boolean + menuLabel?: string + align?: 'left' | 'right' + triggerClassName?: string + title?: string + ariaLabel?: string +} + +/** + * 自选分组选择菜单。 + * 分组仅在菜单打开时读取,React Query 会在多个入口间共享同一份缓存。 + */ +export function WatchlistGroupMenu({ + children, + onSelect, + disabled = false, + preferredGroupId, + includeAll = false, + counts, + total = 0, + disableEmpty = false, + menuLabel = '选择自选分组', + align = 'right', + triggerClassName = '', + title = '加入自选', + ariaLabel = title, +}: WatchlistGroupMenuProps) { + const [open, setOpen] = useState(false) + const [position, setPosition] = useState({ top: 0, left: 0 }) + const triggerRef = useRef(null) + const menuRef = useRef(null) + const menuId = useId() + + const groupsQuery = useQuery({ + queryKey: QK.watchlistGroups, + queryFn: api.watchlistGroups, + enabled: open, + staleTime: 60_000, + }) + const groups = groupsQuery.data?.groups ?? [] + const showPreferred = preferredGroupId !== undefined + + const placeMenu = () => { + const trigger = triggerRef.current + if (!trigger) return + + const rect = trigger.getBoundingClientRect() + const menuHeight = Math.min(menuRef.current?.offsetHeight ?? MENU_MAX_HEIGHT, MENU_MAX_HEIGHT) + const spaceBelow = window.innerHeight - rect.bottom + const spaceAbove = rect.top + const dropUp = spaceBelow < menuHeight + TRIGGER_GAP && spaceAbove > spaceBelow + const top = dropUp + ? Math.max(VIEWPORT_GAP, rect.top - menuHeight - TRIGGER_GAP) + : Math.min(rect.bottom + TRIGGER_GAP, window.innerHeight - menuHeight - VIEWPORT_GAP) + const rawLeft = align === 'left' ? rect.left : rect.right - MENU_WIDTH + const left = Math.max( + VIEWPORT_GAP, + Math.min(rawLeft, window.innerWidth - MENU_WIDTH - VIEWPORT_GAP), + ) + setPosition({ top, left }) + } + + const toggleMenu = () => { + if (disabled) return + if (open) { + setOpen(false) + return + } + placeMenu() + setOpen(true) + } + + useLayoutEffect(() => { + if (!open) return + placeMenu() + }, [open, groups.length, groupsQuery.isPending]) + + useEffect(() => { + if (!open) return + + const closeOnOutsideClick = (event: MouseEvent) => { + const target = event.target as Node + if (triggerRef.current?.contains(target) || menuRef.current?.contains(target)) return + setOpen(false) + } + const closeOnViewportChange = () => setOpen(false) + const closeOnEscape = (event: KeyboardEvent) => { + if (event.key !== 'Escape') return + event.preventDefault() + event.stopPropagation() + setOpen(false) + triggerRef.current?.focus() + } + + document.addEventListener('mousedown', closeOnOutsideClick) + window.addEventListener('keydown', closeOnEscape, true) + window.addEventListener('scroll', closeOnViewportChange, true) + window.addEventListener('resize', closeOnViewportChange) + return () => { + document.removeEventListener('mousedown', closeOnOutsideClick) + window.removeEventListener('keydown', closeOnEscape, true) + window.removeEventListener('scroll', closeOnViewportChange, true) + window.removeEventListener('resize', closeOnViewportChange) + } + }, [open]) + + useEffect(() => { + if (!open || groupsQuery.isPending) return + menuRef.current?.querySelector('[role="menuitem"]:not(:disabled)')?.focus() + }, [open, groups.length, groupsQuery.isPending]) + + const choose = (groupId: string | null) => { + setOpen(false) + onSelect(groupId) + } + + const handleMenuKeyDown = (event: React.KeyboardEvent) => { + if (!['ArrowDown', 'ArrowUp', 'Home', 'End'].includes(event.key)) return + const items = Array.from(menuRef.current?.querySelectorAll('[role="menuitem"]:not(:disabled)') ?? []) + if (items.length === 0) return + + event.preventDefault() + const current = items.indexOf(document.activeElement as HTMLButtonElement) + if (event.key === 'Home') items[0].focus() + else if (event.key === 'End') items[items.length - 1].focus() + else if (event.key === 'ArrowDown') items[(current + 1 + items.length) % items.length].focus() + else items[(current - 1 + items.length) % items.length].focus() + } + + const menuItemClass = 'flex h-8 w-full items-center gap-2 rounded-btn px-2 text-left text-xs text-secondary outline-none transition-colors hover:bg-elevated hover:text-foreground focus:bg-elevated focus:text-foreground disabled:cursor-not-allowed disabled:opacity-40 disabled:hover:bg-transparent disabled:hover:text-secondary' + const showCounts = counts !== undefined + const ungroupedCount = counts?.ungrouped ?? 0 + + return ( + <> + + + {open && createPortal( + diff --git a/frontend/src/lib/api.ts b/frontend/src/lib/api.ts index 7b6f8d9..3f68bcd 100644 --- a/frontend/src/lib/api.ts +++ b/frontend/src/lib/api.ts @@ -200,6 +200,12 @@ export interface MinuteKlineRow { amount: number } +export interface MinuteKlineSession { + date: string + prev_close: number | null + rows: MinuteKlineRow[] +} + export interface PriceLimitInfo { rate: number limit_up: number | null @@ -233,6 +239,27 @@ export interface WatchlistEntry { added_at: string note?: string name?: string | null + group_id?: string | null +} + +export type WatchlistGroupColor = + | 'sky' + | 'blue' + | 'indigo' + | 'violet' + | 'fuchsia' + | 'rose' + | 'orange' + | 'amber' + | 'lime' + | 'emerald' + | 'teal' + | 'cyan' + +export interface WatchlistGroup { + id: string + name: string + color: WatchlistGroupColor } export interface WatchlistImportCandidate { @@ -1398,9 +1425,21 @@ export const api = { source?: 'local' | 'live' | 'none' asset_type?: 'stock' | 'etf' | 'index' price_limit?: PriceLimitInfo | null + prev_close?: number | null }>( `/api/kline/minute?symbol=${encodeURIComponent(symbol)}${date ? `&date=${date}` : ''}`, ), + klineMinuteRange: (symbol: string, days = 10) => + request<{ + symbol: string + name?: string + asset_type: 'stock' | 'etf' | 'index' + requested_days: number + sessions: MinuteKlineSession[] + source: 'local' | 'none' + }>( + `/api/kline/minute-range?symbol=${encodeURIComponent(symbol)}&days=${days}`, + ), indexList: () => request<{ results: IndexInstrument[]; count: number }>('/api/index/list'), indexSearch: (q: string, limit = 20) => request<{ results: IndexInstrument[] }>( @@ -1446,10 +1485,10 @@ export const api = { method: 'POST', body: JSON.stringify({ ...(days ? { days } : {}), ...(extend ? { extend: true } : {}) }), }), - syncMinuteSingle: (symbol: string) => + syncMinuteSingle: (symbol: string, days?: number) => request<{ status: string; symbol: string; rows: number }>('/api/kline/sync_minute_single', { method: 'POST', - body: JSON.stringify({ symbol }), + body: JSON.stringify({ symbol, ...(days != null ? { days } : {}) }), }), clearMinute: () => request<{ status: string; removed: number }>('/api/kline/clear_minute', { @@ -1472,16 +1511,38 @@ export const api = { }), watchlistList: () => request<{ symbols: WatchlistEntry[] }>('/api/watchlist'), - watchlistAdd: (symbol: string, note = '') => + watchlistAdd: (symbol: string, note = '', groupId?: string | null) => request<{ symbols: WatchlistEntry[] }>('/api/watchlist', { method: 'POST', - body: JSON.stringify({ symbol, note }), + body: JSON.stringify({ symbol, note, group_id: groupId ?? null }), }), - watchlistBatchAdd: (symbols: string[], note = '') => + watchlistBatchAdd: (symbols: string[], note = '', groupId?: string | null) => request<{ symbols: WatchlistEntry[]; added: number }>('/api/watchlist/batch', { method: 'POST', - body: JSON.stringify({ symbols, note }), + body: JSON.stringify({ symbols, note, group_id: groupId ?? null }), }), + watchlistGroups: () => + request<{ groups: WatchlistGroup[] }>('/api/watchlist/groups'), + watchlistGroupCreate: (name: string, color: WatchlistGroupColor) => + request<{ groups: WatchlistGroup[]; group: WatchlistGroup }>('/api/watchlist/groups', { + method: 'POST', + body: JSON.stringify({ name, color }), + }), + watchlistGroupRename: (groupId: string, name: string, color: WatchlistGroupColor) => + request<{ groups: WatchlistGroup[] }>( + `/api/watchlist/groups/${encodeURIComponent(groupId)}`, + { method: 'PUT', body: JSON.stringify({ name, color }) }, + ), + watchlistGroupDelete: (groupId: string) => + request<{ groups: WatchlistGroup[]; symbols: WatchlistEntry[] }>( + `/api/watchlist/groups/${encodeURIComponent(groupId)}`, + { method: 'DELETE' }, + ), + watchlistSetGroup: (symbol: string, groupId: string | null) => + request<{ symbols: WatchlistEntry[] }>( + `/api/watchlist/${encodeURIComponent(symbol)}/group`, + { method: 'PUT', body: JSON.stringify({ group_id: groupId }) }, + ), watchlistOcrStatus: () => request<{ provider: string; available: boolean }>('/api/watchlist/ocr-status'), watchlistImportImage: (file: File, signal?: AbortSignal, quiet = false) => { diff --git a/frontend/src/lib/intraday-chart.ts b/frontend/src/lib/intraday-chart.ts new file mode 100644 index 0000000..928b6fd --- /dev/null +++ b/frontend/src/lib/intraday-chart.ts @@ -0,0 +1,40 @@ +import type { MinuteKlineRow } from '@/lib/api' + +export function formatMinuteTime(datetime: string): string { + const match = datetime.match(/(\d{2}):(\d{2})/) + if (!match) return datetime.slice(11, 16) + const hour = (parseInt(match[1]) + 8) % 24 + return `${String(hour).padStart(2, '0')}:${match[2]}` +} + +export function computeIntradayAverage(data: MinuteKlineRow[]): number[] { + const result: number[] = [] + let amount = 0 + let volume = 0 + for (const row of data) { + amount += row.amount + volume += row.volume * 100 + result.push(volume > 0 ? amount / volume : row.close) + } + return result +} + +function generateFullDayTimes(): string[] { + const times: string[] = [] + for (let hour = 9; hour <= 11; hour++) { + const startMinute = hour === 9 ? 30 : 0 + const endMinute = hour === 11 ? 30 : 59 + for (let minute = startMinute; minute <= endMinute; minute++) { + times.push(`${String(hour).padStart(2, '0')}:${String(minute).padStart(2, '0')}`) + } + } + for (let hour = 13; hour <= 15; hour++) { + const endMinute = hour === 15 ? 0 : 59 + for (let minute = 0; minute <= endMinute; minute++) { + times.push(`${String(hour).padStart(2, '0')}:${String(minute).padStart(2, '0')}`) + } + } + return times +} + +export const FULL_DAY_TIMES = generateFullDayTimes() diff --git a/frontend/src/lib/queryKeys.ts b/frontend/src/lib/queryKeys.ts index 50f2e75..27489ac 100644 --- a/frontend/src/lib/queryKeys.ts +++ b/frontend/src/lib/queryKeys.ts @@ -23,6 +23,7 @@ export const QK = { // Watchlist watchlist: ['watchlist'] as const, + watchlistGroups: ['watchlist-groups'] as const, watchlistQuotes: ['watchlist-quotes'] as const, watchlistEnriched: (ext?: string) => ['watchlist-enriched', ext] as const, watchlistKlineBatch: (symbols: string) => ['watchlist-kline-batch', symbols] as const, @@ -61,6 +62,8 @@ export const QK = { stockLevels: (symbol: string, days?: number) => ['stock-levels', symbol, days ?? 120] as const, klineMinute: (symbol: string, date: string) => ['kline-minute', symbol, date] as const, + klineMinuteRange: (symbol: string, days: number) => + ['kline-minute-range', symbol, days] as const, indexDaily: (symbol: string, start: string, end: string) => ['index-daily', symbol, start, end] as const, indexMinute: (symbol: string, date: string) => diff --git a/frontend/src/lib/storage.ts b/frontend/src/lib/storage.ts index 2746d04..1c14271 100644 --- a/frontend/src/lib/storage.ts +++ b/frontend/src/lib/storage.ts @@ -36,6 +36,9 @@ export const storage = { /** 个股日K成交量对比设置 */ stockVolumeCompare: kv<{ enabled: boolean; days: number }>('stock_volume_compare'), + /** 个股详情多日分时周期 */ + stockPreviewIntradayDays: kv('stock_preview_intraday_days'), + /** 策略结果列表列配置 */ screenerResultColumns: kv('screener_result_columns'), diff --git a/frontend/src/lib/useSharedMutations.ts b/frontend/src/lib/useSharedMutations.ts index 0937cba..f825824 100644 --- a/frontend/src/lib/useSharedMutations.ts +++ b/frontend/src/lib/useSharedMutations.ts @@ -29,11 +29,17 @@ export function useUpdateQuoteInterval() { }) } -/** 批量添加自选 — Screener / Intraday / 截图导入 共用 */ +interface WatchlistBatchAddInput { + symbols: string[] + groupId?: string | null +} + +/** 批量添加自选 — Screener / 截图导入共用 */ export function useWatchlistBatchAdd() { const qc = useQueryClient() return useMutation({ - mutationFn: (symbols: string[]) => api.watchlistBatchAdd(symbols), + mutationFn: ({ symbols, groupId }: WatchlistBatchAddInput) => + api.watchlistBatchAdd(symbols, '', groupId), onSuccess: () => { qc.invalidateQueries({ queryKey: QK.watchlist }) // 前缀匹配: 实际 key 为 ['watchlist-enriched', extColumnsParam], diff --git a/frontend/src/lib/watchlist-group-colors.ts b/frontend/src/lib/watchlist-group-colors.ts new file mode 100644 index 0000000..8cde1f7 --- /dev/null +++ b/frontend/src/lib/watchlist-group-colors.ts @@ -0,0 +1,33 @@ +import type { WatchlistGroupColor } from '@/lib/api' + +export interface WatchlistGroupColorOption { + id: WatchlistGroupColor + label: string + text: string + border: string + background: string + dot: string + ring: string +} + +export const DEFAULT_WATCHLIST_GROUP_COLOR: WatchlistGroupColor = 'sky' + +export const WATCHLIST_GROUP_COLORS: readonly WatchlistGroupColorOption[] = [ + { id: 'sky', label: '天蓝', text: 'text-sky-400', border: 'border-sky-400/40', background: 'bg-sky-400/10', dot: 'bg-sky-400', ring: 'ring-sky-400/60' }, + { id: 'blue', label: '蓝色', text: 'text-blue-400', border: 'border-blue-400/40', background: 'bg-blue-400/10', dot: 'bg-blue-400', ring: 'ring-blue-400/60' }, + { id: 'indigo', label: '靛蓝', text: 'text-indigo-400', border: 'border-indigo-400/40', background: 'bg-indigo-400/10', dot: 'bg-indigo-400', ring: 'ring-indigo-400/60' }, + { id: 'violet', label: '紫色', text: 'text-violet-400', border: 'border-violet-400/40', background: 'bg-violet-400/10', dot: 'bg-violet-400', ring: 'ring-violet-400/60' }, + { id: 'fuchsia', label: '品红', text: 'text-fuchsia-400', border: 'border-fuchsia-400/40', background: 'bg-fuchsia-400/10', dot: 'bg-fuchsia-400', ring: 'ring-fuchsia-400/60' }, + { id: 'rose', label: '玫红', text: 'text-rose-400', border: 'border-rose-400/40', background: 'bg-rose-400/10', dot: 'bg-rose-400', ring: 'ring-rose-400/60' }, + { id: 'orange', label: '橙色', text: 'text-orange-400', border: 'border-orange-400/40', background: 'bg-orange-400/10', dot: 'bg-orange-400', ring: 'ring-orange-400/60' }, + { id: 'amber', label: '金色', text: 'text-amber-400', border: 'border-amber-400/40', background: 'bg-amber-400/10', dot: 'bg-amber-400', ring: 'ring-amber-400/60' }, + { id: 'lime', label: '青柠', text: 'text-lime-400', border: 'border-lime-400/40', background: 'bg-lime-400/10', dot: 'bg-lime-400', ring: 'ring-lime-400/60' }, + { id: 'emerald', label: '绿色', text: 'text-emerald-400', border: 'border-emerald-400/40', background: 'bg-emerald-400/10', dot: 'bg-emerald-400', ring: 'ring-emerald-400/60' }, + { id: 'teal', label: '墨绿', text: 'text-teal-400', border: 'border-teal-400/40', background: 'bg-teal-400/10', dot: 'bg-teal-400', ring: 'ring-teal-400/60' }, + { id: 'cyan', label: '青色', text: 'text-cyan-400', border: 'border-cyan-400/40', background: 'bg-cyan-400/10', dot: 'bg-cyan-400', ring: 'ring-cyan-400/60' }, +] + +export function resolveWatchlistGroupColor(color?: string | null): WatchlistGroupColorOption { + return WATCHLIST_GROUP_COLORS.find(option => option.id === color) + ?? WATCHLIST_GROUP_COLORS[0] +} diff --git a/frontend/src/pages/Screener.tsx b/frontend/src/pages/Screener.tsx index c195a63..b6b533f 100644 --- a/frontend/src/pages/Screener.tsx +++ b/frontend/src/pages/Screener.tsx @@ -14,6 +14,7 @@ import { PageHeader } from '@/components/PageHeader' import { EmptyState } from '@/components/EmptyState' import { DatePicker } from '@/components/DatePicker' import { StockPreviewDialog } from '@/components/StockPreviewDialog' +import { WatchlistAddMenu } from '@/components/WatchlistAddMenu' import { useStrategyPool } from '@/lib/useStrategyPool' import { StrategyCard, CardSize, loadCardSize, cardWrapCls } from '@/components/screener/StrategyCard' import { ScreenerTable } from '@/components/screener/ScreenerTable' @@ -509,8 +510,17 @@ export function Screener() { // 单只股票加入/移出自选 const toggleWatchlist = useMutation({ - mutationFn: ({ symbol, inList }: { symbol: string; inList: boolean }) => - inList ? api.watchlistRemove(symbol) : api.watchlistAdd(symbol), + mutationFn: ({ + symbol, + action, + groupId, + }: { + symbol: string + action: 'add' | 'remove' + groupId?: string | null + }) => action === 'remove' + ? api.watchlistRemove(symbol) + : api.watchlistAdd(symbol, '', groupId), onSuccess: () => { qc.invalidateQueries({ queryKey: QK.watchlist }) qc.invalidateQueries({ queryKey: ['watchlist-enriched'] }) @@ -567,10 +577,10 @@ export function Screener() { } } - const handleBatchAdd = () => { + const handleBatchAdd = (groupId: string | null) => { if (!displayRows.length) return const symbols = displayRows.map((r: any) => r.symbol) - batchAdd.mutate(symbols, { + batchAdd.mutate({ symbols, groupId }, { onSuccess: (data) => { setBatchMsg(`已添加 ${data.added} 只到自选`) setTimeout(() => setBatchMsg(''), 3000) @@ -815,16 +825,19 @@ export function Screener() {
)} {displayRows.length > 0 && ( - + )} - + {inWatchlist ? ( + + ) : ( + onAdd(r.symbol, groupId)} + preferredGroupId={preferredGroupId} + disabled={addPending} + triggerClassName="shrink-0 rounded p-1 text-muted transition-colors hover:bg-accent/10 hover:text-accent disabled:opacity-50" + > + + + )}
) })} @@ -402,6 +420,9 @@ const StockCard = React.memo(function StockCard({ onToggleExpand, onDimensionClick, isMonitored, + groups, + onGroupChange, + groupChangePending, }: { r: any candleRows: KlineRow[] @@ -416,6 +437,9 @@ const StockCard = React.memo(function StockCard({ onToggleExpand: (key: string) => void onDimensionClick: (target: DimensionMembersTarget) => void isMonitored?: boolean + groups: WatchlistGroup[] + onGroupChange: (symbol: string, groupId: string | null) => void + groupChangePending: boolean }) { const board = boardTag(r.symbol) const price = r.rt_price ?? r.close @@ -444,7 +468,7 @@ const StockCard = React.memo(function StockCard({ {/* 左侧彩色指示条 */}
- {/* 删除按钮 / 确认区 */} + {/* 分组与删除入口 */}
{isConfirming ? (
e.stopPropagation()}> @@ -459,20 +483,29 @@ const StockCard = React.memo(function StockCard({
) : ( - +
e.stopPropagation()}> + + +
)}
{/* 卡片内容 */}
{/* 第一行: 代码 + 名称 + 板块标识 */} -
+
{r.symbol} @@ -584,6 +617,7 @@ export function Watchlist() { const [columns, setColumns] = useState([...BUILTIN_COLUMNS]) const [customizerOpen, setCustomizerOpen] = useState(false) const [importOpen, setImportOpen] = useState(false) + const [selectedGroup, setSelectedGroup] = useState('all') const [ocrAvailable, setOcrAvailable] = useState(null) const [ocrInstallHint, setOcrInstallHint] = useState('') const columnsLoaded = useRef(false) @@ -696,6 +730,26 @@ export function Watchlist() { queryFn: api.watchlistList, }) + const groupList = useQuery({ + queryKey: QK.watchlistGroups, + queryFn: api.watchlistGroups, + }) + const groups = groupList.data?.groups ?? [] + const activeGroupId = selectedGroup === 'all' || selectedGroup === 'ungrouped' + ? null + : selectedGroup + + useEffect(() => { + if ( + selectedGroup !== 'all' + && selectedGroup !== 'ungrouped' + && groupList.isSuccess + && !groups.some(group => group.id === selectedGroup) + ) { + setSelectedGroup('all') + } + }, [groupList.isSuccess, groups, selectedGroup]) + // enriched 数据 — 传入 ext_columns 参数 const enriched = useQuery({ queryKey: QK.watchlistEnriched(extColumnsParam), @@ -743,7 +797,8 @@ export function Watchlist() { const minuteData = intradayVisible ? (minuteBatch.data?.data ?? {}) : {} const addMutation = useMutation({ - mutationFn: (sym: string) => api.watchlistAdd(sym), + mutationFn: ({ symbol, groupId }: { symbol: string; groupId: string | null }) => + api.watchlistAdd(symbol, '', groupId), onSuccess: (data) => { qc.setQueryData(QK.watchlist, data) qc.invalidateQueries({ queryKey: QK.watchlist }) @@ -791,6 +846,36 @@ export function Watchlist() { }, }) + const createGroup = useMutation({ + mutationFn: ({ name, color }: { name: string; color: WatchlistGroupColor }) => + api.watchlistGroupCreate(name, color), + onSuccess: data => { + qc.setQueryData(QK.watchlistGroups, { groups: data.groups }) + setSelectedGroup(data.group.id) + }, + }) + + const renameGroup = useMutation({ + mutationFn: ({ groupId, name, color }: { groupId: string; name: string; color: WatchlistGroupColor }) => + api.watchlistGroupRename(groupId, name, color), + onSuccess: data => qc.setQueryData(QK.watchlistGroups, data), + }) + + const deleteGroup = useMutation({ + mutationFn: (groupId: string) => api.watchlistGroupDelete(groupId), + onSuccess: (data, groupId) => { + qc.setQueryData(QK.watchlistGroups, { groups: data.groups }) + qc.setQueryData(QK.watchlist, { symbols: data.symbols }) + if (selectedGroup === groupId) setSelectedGroup('all') + }, + }) + + const assignGroup = useMutation({ + mutationFn: ({ symbol, groupId }: { symbol: string; groupId: string | null }) => + api.watchlistSetGroup(symbol, groupId), + onSuccess: data => qc.setQueryData(QK.watchlist, data), + }) + // 二次确认状态 const [confirmClear, setConfirmClear] = useState(false) const [confirmRemove, setConfirmRemove] = useState(null) @@ -804,9 +889,35 @@ export function Watchlist() { }, [remove]) const handleCardCancelRemove = useCallback(() => setConfirmRemove(null), []) const handleCardRequestRemove = useCallback((sym: string) => setConfirmRemove(sym), []) + const handleGroupChange = useCallback((symbol: string, groupId: string | null) => { + assignGroup.mutate({ symbol, groupId }) + }, [assignGroup]) - const allSymbols = list.data?.symbols?.map(s => s.symbol) ?? [] + const listEntries = list.data?.symbols ?? [] + const allSymbols = listEntries.map(s => s.symbol) const rows = enriched.data?.rows ?? [] + const groupBySymbol = useMemo( + () => new Map(listEntries.map(entry => [entry.symbol, entry.group_id ?? null])), + [listEntries], + ) + const groupCounts = useMemo(() => { + const counts: Record = { ungrouped: 0 } + for (const entry of listEntries) { + const groupId = entry.group_id ?? 'ungrouped' + counts[groupId] = (counts[groupId] ?? 0) + 1 + } + return counts + }, [listEntries]) + const rowsInSelectedGroup = useMemo(() => { + const rowsWithGroup = rows.map(row => ({ ...row, group_id: groupBySymbol.get(row.symbol) ?? null })) + if (selectedGroup === 'all') return rowsWithGroup + if (selectedGroup === 'ungrouped') return rowsWithGroup.filter(row => row.group_id == null) + return rowsWithGroup.filter(row => row.group_id === selectedGroup) + }, [groupBySymbol, rows, selectedGroup]) + const activeGroup = activeGroupId + ? groups.find(group => group.id === activeGroupId) + : undefined + const watchlistContentLoading = list.isLoading || (allSymbols.length > 0 && enriched.isLoading) // 实时监控圆点: 仅 Free/低档 "按自选股实时监控" 模式 (mode === 'watchlist') 下显示; // Starter+ 全市场模式 (mode === 'full_market') 全部标的都在监控, 标圆点无意义, 故不显示。 @@ -885,7 +996,7 @@ export function Watchlist() { // 筛选 + 排序 const filteredRows = useMemo(() => { // 板块筛选(全选时跳过) - let result = rows + let result = rowsInSelectedGroup if (boardFilter.size > 0 && boardFilter.size < BOARDS.length) { result = result.filter(r => { // 非股票 (指数/ETF) 无板块语义, 不受板块筛选影响 (顺带修复 ETF 行被误过滤) @@ -914,7 +1025,7 @@ export function Watchlist() { }) } return result - }, [rows, filters, columns, boardFilter]) + }, [rowsInSelectedGroup, filters, columns, boardFilter]) const activeFilterCount = Object.values(filters).filter(v => v.min || v.max || v.text).length const hasBoardFilter = boardFilter.size > 0 && boardFilter.size < BOARDS.length @@ -961,8 +1072,8 @@ export function Watchlist() { ) // "被筛选条件隐藏" 的个股数: 后端返回的行数 vs 经过前端筛选后的行数. - // rows.length 是后端实际返回 (含 pending 行), 减去 sortedRows (筛选后) 才是真正的筛选隐藏. - const hiddenCount = Math.max(0, rows.length - sortedRows.length) + // 分组切换不计入筛选隐藏,只比较当前分组内的数据。 + const hiddenCount = Math.max(0, rowsInSelectedGroup.length - sortedRows.length) const renderStockCard = (r: any) => ( ) @@ -993,7 +1107,7 @@ export function Watchlist() { {sortedRows.length} / - {allSymbols.length} + {rowsInSelectedGroup.length} {/* 数据未就绪提示: 自选了但 enriched 缓存未覆盖 (新股/冷门/新用户未同步), 指标全为 null */} @@ -1045,7 +1159,9 @@ export function Watchlist() { { setPreviewSymbol(sym); setPreviewName(name) }} existingSymbols={allSymbols as string[]} - onAdd={(sym) => addMutation.mutate(sym)} + onAdd={(symbol, groupId) => addMutation.mutate({ symbol, groupId })} + preferredGroupId={activeGroupId} + addPending={addMutation.isPending} />
) } diff --git a/frontend/src/pages/backtest/StrategyBacktest.tsx b/frontend/src/pages/backtest/StrategyBacktest.tsx index 58f966f..8f1d4b0 100644 --- a/frontend/src/pages/backtest/StrategyBacktest.tsx +++ b/frontend/src/pages/backtest/StrategyBacktest.tsx @@ -29,6 +29,7 @@ import { StrategyNavChart } from './charts/StrategyNavChart' import { ReturnDistributionChart } from './charts/ReturnDistributionChart' import { TradeKlineModal } from './components/TradeKlineModal' import { SignalTriggerActions } from '@/components/signals/SignalTriggerActions' +import { WatchlistGroupMenu } from '@/components/WatchlistAddMenu' const formatDate = (date: Date) => date.toISOString().slice(0, 10) const monthsAgo = (months: number) => { @@ -774,6 +775,15 @@ function StockPoolPicker({ value, onChange, assetType = 'stock' }: { value: stri queryFn: () => api.watchlistList(), staleTime: 30_000, }) + const watchlistEntries = watchlist.data?.symbols ?? [] + const watchlistCounts = useMemo(() => { + const counts: Record = { ungrouped: 0 } + for (const entry of watchlistEntries) { + const groupId = entry.group_id ?? 'ungrouped' + counts[groupId] = (counts[groupId] ?? 0) + 1 + } + return counts + }, [watchlistEntries]) useEffect(() => { if (results.length === 0) return @@ -802,9 +812,11 @@ function StockPoolPicker({ value, onChange, assetType = 'stock' }: { value: stri setOpen(false) } const removeSymbol = (symbol: string) => setSymbols(symbols.filter(s => s !== symbol)) - // 一键导入自选: 合并去重, 顺带回填股票名 - const importFromWatchlist = () => { - const entries = watchlist.data?.symbols ?? [] + // 按分组导入自选: 合并去重, 顺带回填股票名 + const importFromWatchlist = (groupId: string | null) => { + const entries = groupId === 'all' + ? watchlistEntries + : watchlistEntries.filter(entry => (entry.group_id ?? null) === groupId) if (entries.length === 0) return setSymbolNames(prev => { const next = { ...prev } @@ -813,7 +825,7 @@ function StockPoolPicker({ value, onChange, assetType = 'stock' }: { value: stri }) setSymbols([...symbols, ...entries.map(e => e.symbol)]) } - const watchlistCount = watchlist.data?.symbols?.length ?? 0 + const watchlistCount = watchlistEntries.length return (
@@ -861,16 +873,22 @@ function StockPoolPicker({ value, onChange, assetType = 'stock' }: { value: stri {symbols.length === 0 ? '全市场' : `共 ${symbols.length} 只`} - + +
+ + {/* 状态卡 — 收起时隐藏 */} + {!navCollapsed && ( +
+ +
-
- -
- Quant · Terminal -
- -
- - - + )}
- {/* 数据源状态条 */} -
- - - {/* 全局行情开关 */} + ) : (
{isNoneTier && !realtimeProviderName ? (
@@ -634,28 +717,41 @@ export function Layout() { )}
+ )} -
-
+
+
cn( - 'flex flex-1 items-center justify-between gap-3 px-3 py-2 rounded-btn text-sm transition-colors duration-150 ease-smooth', + 'group relative flex items-center rounded-btn text-sm transition-all duration-150 ease-smooth', + navCollapsed ? 'justify-center px-0 py-2' : 'flex-1 gap-3 px-3 py-2', isActive ? 'bg-elevated text-foreground font-medium' - : 'text-foreground/80 hover:bg-elevated hover:text-foreground', + : 'text-foreground/75 hover:bg-elevated/70 hover:text-foreground', ) } > - - - 设置 - - - {version ?? ''} - + {({ isActive }) => ( + <> + + + {!navCollapsed && 设置} + {!navCollapsed && version && ( + + {version} + + )} + + )}
diff --git a/frontend/src/components/StockPanel.tsx b/frontend/src/components/StockPanel.tsx index 4e4f707..2da0476 100644 --- a/frontend/src/components/StockPanel.tsx +++ b/frontend/src/components/StockPanel.tsx @@ -1,4 +1,5 @@ import { useEffect, useState, useCallback, useRef, useMemo } from 'react' +import { X } from 'lucide-react' import { type KlineRow, type FinancialMetricRecord } from '@/lib/api' import { StockInfoBar } from '@/components/StockInfoBar' import { StockDailyKChart, getDefaultRange, type StockDailyKChartResult } from '@/components/StockDailyKChart' @@ -61,6 +62,7 @@ export function StockPanel({ }: Props) { const [linkedPrice, setLinkedPrice] = useState(null) const [selectedDate, setSelectedDate] = useState(null) + const [intradayDismissed, setIntradayDismissed] = useState(false) const [dailyResult, setDailyResult] = useState(null) // 信息条指标配置提升到此层:同时供 StockInfoBar 渲染与 StockDailyKChart 请求 ext 数据 const [fields, setFields] = useState(loadInfoFields) @@ -86,6 +88,7 @@ export function StockPanel({ const handleDateClick = useCallback((date: string) => { setSelectedDate(date) + setIntradayDismissed(false) onSelectDate?.(date) }, [onSelectDate]) @@ -159,16 +162,25 @@ export function StockPanel({ extColumns={extColumns} /> - {showIntraday && selectedDate && ( - + {showIntraday && selectedDate && !intradayDismissed && ( +
+ + +
)}
diff --git a/frontend/src/components/StockPreviewDialog.tsx b/frontend/src/components/StockPreviewDialog.tsx index 874b189..828a30e 100644 --- a/frontend/src/components/StockPreviewDialog.tsx +++ b/frontend/src/components/StockPreviewDialog.tsx @@ -1,11 +1,13 @@ import { useState, useEffect } from 'react' import { useQuery, useMutation, useQueryClient } from '@tanstack/react-query' import { motion, AnimatePresence } from 'framer-motion' -import { X, RefreshCw, Clock, LineChart } from 'lucide-react' +import { X, RefreshCw, Clock, LineChart, Star, RadioTower, Maximize2, Minimize2 } from 'lucide-react' import { api } from '@/lib/api' import { QK } from '@/lib/queryKeys' +import { cn } from '@/lib/cn' import { cnSignal } from '@/lib/signals' import { StockPanel, getDefaultRange } from '@/components/StockPanel' +import { WatchlistAddMenu } from '@/components/WatchlistAddMenu' import { StockMultiDayIntradayChart } from '@/components/StockMultiDayIntradayChart' import { DatePicker } from '@/components/DatePicker' import { RuleEditor } from '@/components/monitor/RuleEditor' @@ -58,6 +60,7 @@ export function StockPreviewDialog({ symbol, name, onClose, triggerInfo }: Props const [intradayDays, setIntradayDays] = useState(loadIntradayDays) const [dateRange, setDateRange] = useState(getDefaultRange) const [showMonitorEditor, setShowMonitorEditor] = useState(false) + const [maximized, setMaximized] = useState(false) const qc = useQueryClient() const backdrop = useDialogBackdrop(onClose) @@ -152,7 +155,10 @@ export function StockPreviewDialog({ symbol, name, onClose, triggerInfo }: Props animate={{ opacity: 1, scale: 1, y: 0 }} exit={{ opacity: 0, scale: 0.97, y: 8 }} transition={{ duration: 0.2, ease: [0.16, 1, 0.3, 1] }} - className="relative w-[92vw] max-w-[1100px] max-h-[95vh] rounded-card border border-border bg-base shadow-2xl overflow-hidden flex flex-col" + className={cn( + 'relative rounded-card border border-border bg-base shadow-2xl overflow-hidden flex flex-col transition-all duration-200 ease-smooth', + maximized ? 'w-screen h-screen max-w-none max-h-none' : 'w-[92vw] max-w-[1100px] max-h-[95vh]', + )} > {/* 顶栏 */}
@@ -169,88 +175,49 @@ export function StockPreviewDialog({ symbol, name, onClose, triggerInfo }: Props {name && {name}}
- -
- -
-
- - -
- -
+
+ {/* 区间选择 — 随视图切换 */} {view === 'daily' ? ( - <> - {/* 日期范围快捷 */} - {PRESETS.map(p => { - const now = new Date() - const s = new Date(now) - s.setMonth(s.getMonth() - p.months) - const expected = s.toISOString().slice(0, 10) - const isActive = dateRange.start === expected - return ( - - ) - })} - setDateRange(prev => ({ ...prev, start: v }))} - max={dateRange.end} - /> - ~ - setDateRange(prev => ({ ...prev, end: v }))} - min={dateRange.start} - /> - +
+ {PRESETS.map(p => { + const now = new Date() + const s = new Date(now) + s.setMonth(s.getMonth() - p.months) + const expected = s.toISOString().slice(0, 10) + const isActive = dateRange.start === expected + return ( + + ) + })} + setDateRange(prev => ({ ...prev, start: v }))} + max={dateRange.end} + /> + ~ + setDateRange(prev => ({ ...prev, end: v }))} + min={dateRange.start} + /> +
) : ( - <> - 区间 +
{INTRADAY_DAY_OPTIONS.map(days => (
- +
)} + {/* 日K / 分时 切换 */} +
+ + +
+ + + + {/* 自选 */} + {inWatchlist ? ( + + ) : ( + toggleWatchlist.mutate({ action: 'add', groupId })} + disabled={toggleWatchlist.isPending} + triggerClassName="rounded-btn p-1.5 text-muted transition-colors cursor-pointer hover:bg-elevated hover:text-foreground disabled:opacity-50" + ariaLabel={`将 ${symbol} 加入自选`} + > + + + )} + {/* 加监控 */} + + {/* 刷新 */} + + {/* 放大 / 缩小 */} + + +
@@ -331,13 +377,8 @@ export function StockPreviewDialog({ symbol, name, onClose, triggerInfo }: Props setShowMonitorEditor(true)} - inWatchlist={inWatchlist} - onAddToWatchlist={groupId => toggleWatchlist.mutate({ action: 'add', groupId })} - onRemoveFromWatchlist={() => toggleWatchlist.mutate({ action: 'remove' })} - watchlistPending={toggleWatchlist.isPending} /> ) : ( Promise onRename: (groupId: string, name: string, color: WatchlistGroupColor) => Promise onDelete: (groupId: string) => Promise + onClearGroup?: (groupId: string) => Promise } export function WatchlistGroupBar({ @@ -30,8 +35,10 @@ export function WatchlistGroupBar({ onCreate, onRename, onDelete, + onClearGroup, }: GroupBarProps) { const [managerOpen, setManagerOpen] = useState(false) + const [confirmClear, setConfirmClear] = useState(false) const tabs = [ { id: 'all', name: '全部', count: total, color: null }, { id: 'ungrouped', name: '未分组', count: counts.ungrouped ?? 0, color: null }, @@ -80,8 +87,50 @@ export function WatchlistGroupBar({ > + {/* 清空当前分组 — 仅选中具体分组时显示 */} + {onClearGroup && selected !== 'all' && selected !== 'ungrouped' && ( + + )}
+ {/* 清空分组确认弹窗 */} + {confirmClear && selected !== 'all' && selected !== 'ungrouped' && ( +
+
setConfirmClear(false)} + /> +
+

清空分组

+

+ 确认清空「{tabs.find(t => t.id === selected)?.name}」分组? 分组内所有股票将转为未分组(不从自选中删除)。 +

+
+ + +
+
+
+ )} + {managerOpen && ( & { onClose: () => void }) { +}: Omit & { onClose: () => void }) { const inputRef = useRef(null) const [newName, setNewName] = useState('') const [newColor, setNewColor] = useState(DEFAULT_WATCHLIST_GROUP_COLOR) @@ -148,6 +197,21 @@ function GroupManagerDialog({ const [pending, setPending] = useState(false) const [error, setError] = useState('') + // 「显示在侧边栏」偏好开关 + const qc = useQueryClient() + const prefs = usePreferences() + const groupsInNav = prefs.data?.watchlist_groups_in_nav ?? false + const [navTogglePending, setNavTogglePending] = useState(false) + const toggleGroupsInNav = async (enabled: boolean) => { + setNavTogglePending(true) + try { + await api.updateWatchlistGroupsInNav(enabled) + await qc.invalidateQueries({ queryKey: QK.preferences }) + } finally { + setNavTogglePending(false) + } + } + const validate = (name: string) => { const value = name.trim() if (!value) return '请输入分组名称' @@ -210,6 +274,27 @@ function GroupManagerDialog({
+ {/* 显示在侧边栏 开关 */} +
+
+
显示在侧边栏
+
开启后可在左侧菜单展开分组子菜单
+
+ +
+
+
+
+
拉取起始时间 (留空=不限)
+ setTimeWindowStart(e.target.value)} + className="w-full rounded-btn border border-border bg-elevated px-2 py-1.5 text-[10px] font-mono text-foreground" + /> +
+
+
拉取结束时间 (留空=不限)
+ setTimeWindowEnd(e.target.value)} + className="w-full rounded-btn border border-border bg-elevated px-2 py-1.5 text-[10px] font-mono text-foreground" + /> +
+
+
字段映射 (外部名 → 内部名,JSON,可选)