fix: restore polished account experience and align AI with Lark design

This commit is contained in:
yuxuanhui
2026-09-07 23:31:23 +08:00
parent 3f3d417296
commit 7dcf46836f
23 changed files with 1669 additions and 807 deletions
+65
View File
@@ -0,0 +1,65 @@
"""Account snapshots with an explicit local daily simulation allowance."""
from datetime import datetime
from zoneinfo import ZoneInfo
from .alphas import number
LOCAL_DAILY_SIMULATION_LIMIT = 10_000
def daily_activity(raw, today, daily_limit=None):
"""Read recordset columns by schema; missing dates are not reported as zero usage."""
recordset = raw.get("records", {})
if not isinstance(recordset, dict):
raise ValueError("Invalid activity recordset")
schema = recordset.get("schema", {})
if not isinstance(schema, dict):
raise ValueError("Invalid activity schema")
properties, rows = schema.get("properties", []), recordset.get("records", [])
if not isinstance(properties, list) or not isinstance(rows, list):
raise ValueError("Invalid activity records")
names = [p.get("name") if isinstance(p, dict) else p for p in properties]
count = None
if "date" in names and "value" in names:
for row in rows:
if isinstance(row, list) and len(row) >= len(names) and row[names.index("date")] == today:
count = number(row[names.index("value")])
result = {
"today": count,
"limit": daily_limit,
"remaining": max(0, daily_limit - count) if daily_limit is not None and count is not None else None,
}
for key in ("yesterday", "total"):
period = raw.get(key)
result[key] = number(period.get("value")) if isinstance(period, dict) else None
if key == "yesterday":
result["yesterday_date"] = period.get("end") if isinstance(period, dict) else None
# The allowance is a user-chosen local budget, not a platform quota.
# Missing activity must not imply that the whole daily budget is available.
return result
def usage_snapshot(data, errors, at=None):
today = (at or datetime.now(ZoneInfo("America/New_York"))).date().isoformat()
activities, errors = {}, dict(errors)
for key in ("simulations", "submissions"):
if key in data:
try:
activities[key] = daily_activity(
data[key],
today,
daily_limit=LOCAL_DAILY_SIMULATION_LIMIT if key == "simulations" else None,
)
except ValueError:
errors[key] = "平台用量数据格式无法识别"
return {
"date": today,
"timezone": "America/New_York",
**activities,
"alphas": {
key: number(data.get("alphas", {}).get(key))
for key in ("unsubmitted", "active", "decommissioned")
},
"errors": errors,
}
+15
View File
@@ -159,12 +159,16 @@ class Runner:
async def refresh_profile(self):
raw = await self.client.profile()
# Confirm identity before fetching or storing additional account data.
async with self.sessions() as db:
account = await db.get(Account, 1)
user_id = raw.get("id")
if not user_id or (account.wq_user_id and account.wq_user_id != str(user_id)):
self.client.disconnect()
raise WqError("平台账户身份不匹配,请核对凭据", "identity_mismatch")
usage = await self.client.account_usage()
async with self.sessions() as db:
account = await db.get(Account, 1)
account.wq_user_id = str(user_id)
# An allowlist avoids persisting unknown personal/security fields.
account.profile = sanitize(
@@ -177,6 +181,14 @@ class Runner:
"email",
"firstName",
"lastName",
"fullName",
"level",
"geniusLevel",
"verified",
"approved",
"dateCreated",
"dateVerified",
"dateApproved",
"role",
"roles",
"permissions",
@@ -185,6 +197,9 @@ class Runner:
if k in raw
}
)
if self.client.permissions is not None:
account.profile = {**account.profile, "permissions": self.client.permissions}
account.profile = {**account.profile, "usage": usage}
account.connection_status, account.connection_error = "connected", None
account.verification_url, account.last_synced_at = None, now()
await db.execute(
+9 -5
View File
@@ -45,7 +45,7 @@ from .schemas import (
from .security import bootstrap, cipher, issue_session, require_auth, token_hash, valid_password
def account_output(account):
def account_output(account, client):
keys = (
"email",
"wq_user_id",
@@ -59,7 +59,11 @@ def account_output(account):
"timezone",
"page_size",
)
return {**{k: getattr(account, k) for k in keys}, "configured": bool(account.password_encrypted)}
return {
**{k: getattr(account, k) for k in keys},
"configured": bool(account.password_encrypted),
"session": client.session_info(),
}
def csv_cell(value):
@@ -189,7 +193,7 @@ def create_app(settings=None, wq_client=None, ai_model_factory=None):
@api.get("/account", response_model=AccountOutput, tags=["account"])
async def get_account():
async with sessions() as db:
return account_output(await db.get(Account, 1))
return account_output(await db.get(Account, 1), runner.client)
@api.put("/account/credentials", response_model=AccountOutput, tags=["account"])
async def credentials(body: CredentialsInput):
@@ -203,7 +207,7 @@ def create_app(settings=None, wq_client=None, ai_model_factory=None):
account.email = body.email
account.password_encrypted = cipher(settings).encrypt(body.password.encode()).decode()
await db.commit()
return account_output(account)
return account_output(account, runner.client)
@api.patch("/account/preferences", response_model=AccountOutput, tags=["account"])
async def preferences(body: PreferencesInput):
@@ -212,7 +216,7 @@ def create_app(settings=None, wq_client=None, ai_model_factory=None):
for key, value in body.model_dump().items():
setattr(account, key, value)
await db.commit()
return account_output(account)
return account_output(account, runner.client)
async def account_job(kind):
async with sessions() as db:
+8
View File
@@ -232,6 +232,13 @@ class JobOutput(BaseModel):
updated_at: datetime
class PlatformSessionOutput(BaseModel):
authenticated: bool
expires_at: datetime | None
remaining_seconds: int | None
total_seconds: float | None
class AccountOutput(BaseModel):
email: str | None
configured: bool
@@ -245,6 +252,7 @@ class AccountOutput(BaseModel):
theme: Literal["light", "dark"]
timezone: str
page_size: int
session: PlatformSessionOutput
class SessionOutput(BaseModel):
+71 -3
View File
@@ -7,7 +7,7 @@ contain credentials, cookies, or temporary authentication links.
import asyncio
import math
import random
from datetime import datetime, timezone
from datetime import datetime, timedelta, timezone
from email.utils import parsedate_to_datetime
from typing import Awaitable, Callable
from urllib.parse import urljoin, urlparse
@@ -41,6 +41,9 @@ class WqClient:
self.auth_generation = 0
self.credentials: tuple[str, str] | None = None
self.verification_url: str | None = None
self.permissions: list[str] | None = None
self.session_expires_at: datetime | None = None
self.session_duration: float | None = None
self.sleep = sleep
self.on_retry: Callable[[float], Awaitable[None]] | None = None
@@ -51,8 +54,47 @@ class WqClient:
self.authenticated = False
self.credentials = None
self.verification_url = None
self.permissions = None
self.session_expires_at = None
self.session_duration = None
self.client.cookies.clear()
def capture_session(self, response: httpx.Response):
"""Keep only permissions and expiry metadata; never retain the authentication body."""
self.permissions, self.session_expires_at, self.session_duration = None, None, None
try:
data = response.json()
except ValueError:
return
if not isinstance(data, dict):
return
permissions = data.get("permissions")
if isinstance(permissions, list) and all(isinstance(p, str) for p in permissions):
self.permissions = list(dict.fromkeys(permissions))
token = data.get("token")
expiry = token.get("expiry") if isinstance(token, dict) else None
if (
isinstance(expiry, (int, float))
and not isinstance(expiry, bool)
and math.isfinite(expiry)
and 0 <= expiry <= 31536000
):
self.session_duration = expiry
self.session_expires_at = datetime.now(timezone.utc) + timedelta(seconds=expiry)
def session_info(self):
remaining = (
max(0, int((self.session_expires_at - datetime.now(timezone.utc)).total_seconds()))
if self.session_expires_at
else None
)
return {
"authenticated": self.authenticated and remaining != 0,
"expires_at": self.session_expires_at,
"remaining_seconds": remaining,
"total_seconds": self.session_duration,
}
def safe_verification_url(self, response: httpx.Response) -> str:
url = urljoin(str(response.url), response.headers.get("Location", ""))
base = urlparse(self.settings.wq_base_url)
@@ -111,6 +153,7 @@ class WqClient:
raise VerificationRequired(self.verification_url)
self.credentials = (email, password)
self.authenticated = False
self.permissions, self.session_expires_at, self.session_duration = None, None, None
if force:
self.client.cookies.clear()
self.verification_url = None
@@ -126,6 +169,7 @@ class WqClient:
if response.status_code not in (200, 201):
raise WqError(f"平台认证返回异常状态 {response.status_code}", "authentication_failed")
self.authenticated = True
self.capture_session(response)
self.auth_generation += 1
self.verification_url = None
@@ -139,6 +183,7 @@ class WqClient:
response = await self._request("POST", self.verification_url)
if response.status_code in (200, 201):
self.authenticated = True
self.capture_session(response)
self.auth_generation += 1
self.verification_url = None
elif response.status_code in (401, 403, 202):
@@ -146,7 +191,7 @@ class WqClient:
else:
raise WqError("验证会话已失效,请重新连接", "authentication_failed")
async def get(self, path: str, params=None):
async def get(self, path: str, params=None, headers=None):
if not self.credentials:
raise WqError("请先连接 WorldQuant", "disconnected")
if not self.authenticated:
@@ -154,7 +199,7 @@ class WqClient:
refreshed = False
for attempt in range(self.settings.retry_attempts):
generation = self.auth_generation
response = await self._request("GET", path, params=params)
response = await self._request("GET", path, params=params, headers=headers)
if response.status_code == 401 and not refreshed:
await self.authenticate(*self.credentials, stale_generation=generation)
refreshed = True
@@ -189,6 +234,29 @@ class WqClient:
async def profile(self):
return await self.get("/users/self")
async def account_usage(self):
"""Read independent account resources; unavailable sections do not hide the profile."""
from .account_data import usage_snapshot
resources = {
"simulations": "/users/self/activities/simulations",
"submissions": "/users/self/activities/submissions",
"alphas": "/users/self/alphas/summary",
}
data, errors = {}, {}
for key, path in resources.items():
try:
data[key] = await self.get(
path, headers={"Accept": "application/json;version=4.0"} if key == "alphas" else None
)
except VerificationRequired:
raise
except WqError as exc:
if exc.code in ("authentication_failed", "disconnected"):
raise
errors[key] = str(exc)
return usage_snapshot(data, errors)
async def alphas(self, submission, hidden, offset, before):
# Cover every platform stage; submitted records are not assumed to be OS only.
status_key = "status" if submission == "UNSUBMITTED" else "status!"