view schwab_trader/schwab_client.py @ 265:056790c4fb0d

add role-aware Epi assistant prompts Add verified June knowledge, guest/member/admin Copilot profiles, profile-isolated session recovery, animated Epi greetings, and a single authoritative runtime config workflow for inference. Co-authored-by: Copilot <[email protected]>
author MrJuneJune <me@mrjunejune.com>
date Fri, 07 Aug 2026 10:50:30 -0700
parents eb8b4230fdb9
children
line wrap: on
line source

from __future__ import annotations

import base64
import json
import os
import stat
import time
import urllib.error
import urllib.parse
import urllib.request
from dataclasses import dataclass
from pathlib import Path
from typing import Any


AUTH_URL = "https://api.schwabapi.com/v1/oauth/authorize"
TOKEN_URL = "https://api.schwabapi.com/v1/oauth/token"
TRADER_BASE_URL = "https://api.schwabapi.com/trader/v1"
DEFAULT_TOKEN_FILE = "~/.config/zenbu/schwab_tokens.json"


class SchwabError(RuntimeError):
    pass


@dataclass(frozen=True)
class SchwabConfig:
    app_key: str
    app_secret: str
    redirect_uri: str
    token_file: Path

    @classmethod
    def from_env(cls) -> "SchwabConfig":
        app_key = os.environ.get("SCHWAB_APP_KEY", "").strip()
        app_secret = os.environ.get("SCHWAB_APP_SECRET", "").strip()
        redirect_uri = os.environ.get("SCHWAB_REDIRECT_URI", "").strip()
        token_file = Path(os.environ.get("SCHWAB_TOKEN_FILE", DEFAULT_TOKEN_FILE)).expanduser()

        missing = [
            name
            for name, value in (
                ("SCHWAB_APP_KEY", app_key),
                ("SCHWAB_APP_SECRET", app_secret),
                ("SCHWAB_REDIRECT_URI", redirect_uri),
            )
            if not value
        ]
        if missing:
            raise SchwabError("Missing required environment variables: " + ", ".join(missing))

        return cls(
            app_key=app_key,
            app_secret=app_secret,
            redirect_uri=redirect_uri,
            token_file=token_file,
        )


@dataclass(frozen=True)
class ApiResponse:
    status: int
    headers: dict[str, str]
    body: Any
    raw_body: str


def build_authorization_url(app_key: str, redirect_uri: str, state: str | None = None) -> str:
    params = {
        "response_type": "code",
        "client_id": app_key,
        "redirect_uri": redirect_uri,
    }
    if state:
        params["state"] = state
    return AUTH_URL + "?" + urllib.parse.urlencode(params)


def extract_authorization_code(code_or_url: str) -> str:
    value = code_or_url.strip()
    if not value:
        raise SchwabError("Authorization code is empty")

    parsed = urllib.parse.urlparse(value)
    if parsed.scheme and parsed.netloc:
        query = urllib.parse.parse_qs(parsed.query)
        codes = query.get("code")
        if not codes or not codes[0]:
            raise SchwabError("No code= parameter found in callback URL")
        return codes[0]

    return urllib.parse.unquote(value)


def exchange_code_for_tokens(config: SchwabConfig, code_or_url: str) -> dict[str, Any]:
    code = extract_authorization_code(code_or_url)
    return _token_request(
        config,
        {
            "grant_type": "authorization_code",
            "code": code,
            "redirect_uri": config.redirect_uri,
        },
    )


def refresh_tokens(config: SchwabConfig, refresh_token: str | None = None) -> dict[str, Any]:
    token_value = refresh_token
    if token_value is None:
        existing = load_tokens(config.token_file)
        token_value = existing.get("refresh_token")
    if not token_value:
        raise SchwabError("No refresh token available")

    return _token_request(
        config,
        {
            "grant_type": "refresh_token",
            "refresh_token": token_value,
        },
    )


def save_tokens(path: Path, tokens: dict[str, Any]) -> None:
    payload = dict(tokens)
    payload["saved_at"] = int(time.time())
    path.parent.mkdir(parents=True, exist_ok=True)
    path.write_text(json.dumps(payload, indent=2, sort_keys=True) + "\n", encoding="utf-8")
    path.chmod(stat.S_IRUSR | stat.S_IWUSR)


def load_tokens(path: Path) -> dict[str, Any]:
    if not path.exists():
        raise SchwabError(f"Token file does not exist: {path}")
    return json.loads(path.read_text(encoding="utf-8"))


def get_account_numbers(access_token: str) -> ApiResponse:
    return _api_request("GET", "/accounts/accountNumbers", access_token)


def get_accounts(access_token: str, fields: str | None = None) -> ApiResponse:
    query = ""
    if fields:
        query = "?" + urllib.parse.urlencode({"fields": fields})
    return _api_request("GET", "/accounts" + query, access_token)


def get_account(access_token: str, account_hash: str, fields: str | None = None) -> ApiResponse:
    query = ""
    if fields:
        query = "?" + urllib.parse.urlencode({"fields": fields})
    return _api_request("GET", f"/accounts/{urllib.parse.quote(account_hash)}{query}", access_token)


def build_equity_order(
    action: str,
    symbol: str,
    quantity: float,
    order_type: str = "MARKET",
    price: float | None = None,
    duration: str = "DAY",
    session: str = "NORMAL",
) -> dict[str, Any]:
    normalized_action = action.upper()
    normalized_symbol = symbol.upper()
    normalized_order_type = order_type.upper()
    normalized_duration = duration.upper()
    normalized_session = session.upper()

    if normalized_action not in {"BUY", "SELL"}:
        raise SchwabError("action must be BUY or SELL")
    if not normalized_symbol:
        raise SchwabError("symbol is required")
    if quantity <= 0:
        raise SchwabError("quantity must be greater than zero")
    if normalized_order_type not in {"MARKET", "LIMIT"}:
        raise SchwabError("order_type must be MARKET or LIMIT")
    if normalized_order_type == "LIMIT" and price is None:
        raise SchwabError("LIMIT orders require --price")
    if normalized_order_type == "MARKET" and price is not None:
        raise SchwabError("MARKET orders cannot include --price")

    order: dict[str, Any] = {
        "orderType": normalized_order_type,
        "session": normalized_session,
        "duration": normalized_duration,
        "orderStrategyType": "SINGLE",
        "orderLegCollection": [
            {
                "instruction": normalized_action,
                "quantity": quantity,
                "instrument": {
                    "symbol": normalized_symbol,
                    "assetType": "EQUITY",
                },
            }
        ],
    }
    if price is not None:
        order["price"] = f"{price:.2f}"
    return order


def place_order(access_token: str, account_hash: str, order: dict[str, Any]) -> ApiResponse:
    path = f"/accounts/{urllib.parse.quote(account_hash)}/orders"
    return _api_request("POST", path, access_token, order)


def _token_request(config: SchwabConfig, form: dict[str, str]) -> dict[str, Any]:
    credentials = f"{config.app_key}:{config.app_secret}".encode("utf-8")
    headers = {
        "Authorization": "Basic " + base64.b64encode(credentials).decode("ascii"),
        "Content-Type": "application/x-www-form-urlencoded",
        "Accept": "application/json",
    }
    data = urllib.parse.urlencode(form).encode("utf-8")
    request = urllib.request.Request(TOKEN_URL, data=data, headers=headers, method="POST")
    response = _open_request(request)
    if not isinstance(response.body, dict):
        raise SchwabError("Token endpoint did not return a JSON object")
    return response.body


def _api_request(
    method: str,
    path: str,
    access_token: str,
    body: dict[str, Any] | None = None,
) -> ApiResponse:
    headers = {
        "Authorization": f"Bearer {access_token}",
        "Accept": "application/json",
    }
    data = None
    if body is not None:
        data = json.dumps(body).encode("utf-8")
        headers["Content-Type"] = "application/json"

    request = urllib.request.Request(
        TRADER_BASE_URL + path,
        data=data,
        headers=headers,
        method=method,
    )
    return _open_request(request)


def _open_request(request: urllib.request.Request) -> ApiResponse:
    try:
        with urllib.request.urlopen(request, timeout=30) as response:
            raw = response.read().decode("utf-8")
            return ApiResponse(
                status=response.status,
                headers=dict(response.headers.items()),
                body=_parse_json_or_text(raw),
                raw_body=raw,
            )
    except urllib.error.HTTPError as error:
        raw = error.read().decode("utf-8", errors="replace")
        raise SchwabError(
            f"Schwab API request failed with HTTP {error.code}: {raw or error.reason}"
        ) from error
    except urllib.error.URLError as error:
        raise SchwabError(f"Schwab API request failed: {error.reason}") from error


def _parse_json_or_text(raw: str) -> Any:
    if not raw:
        return None
    try:
        return json.loads(raw)
    except json.JSONDecodeError:
        return raw