fix: restore polished account experience and align AI with Lark design
This commit is contained in:
@@ -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,
|
||||
}
|
||||
@@ -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
@@ -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:
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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!"
|
||||
|
||||
Reference in New Issue
Block a user