first commit
This commit is contained in:
@@ -0,0 +1,10 @@
|
||||
import { defineConfig } from 'drizzle-kit'
|
||||
|
||||
export default defineConfig({
|
||||
schema: './src/db/schema.ts',
|
||||
out: './drizzle',
|
||||
dialect: 'sqlite',
|
||||
dbCredentials: {
|
||||
url: process.env.CSBIE_DATABASE_PATH ?? './data/csbie.sqlite',
|
||||
},
|
||||
})
|
||||
@@ -0,0 +1,31 @@
|
||||
{
|
||||
"name": "@repo/csbie-server",
|
||||
"private": true,
|
||||
"type": "module",
|
||||
"exports": {
|
||||
"./app": "./src/app.ts",
|
||||
"./config": "./src/config.ts",
|
||||
"./db": "./src/db/index.ts"
|
||||
},
|
||||
"scripts": {
|
||||
"dev": "bun --env-file=../../.env --watch src/index.ts",
|
||||
"clean": "rm -rf dist drizzle",
|
||||
"typecheck": "tsc",
|
||||
"db:push": "bun --env-file=../../.env drizzle-kit push"
|
||||
},
|
||||
"dependencies": {
|
||||
"@hono/mcp": "^0.3.0",
|
||||
"@modelcontextprotocol/sdk": "^1.29.0",
|
||||
"@napi-rs/keyring": "^1.2.0",
|
||||
"@repo/sbi-client": "workspace:*",
|
||||
"@simplewebauthn/server": "^13.1.2",
|
||||
"drizzle-orm": "^0.44.2",
|
||||
"hono": "^4.8.3",
|
||||
"zod": "^4.4.3"
|
||||
},
|
||||
"devDependencies": {
|
||||
"@types/bun": "latest",
|
||||
"drizzle-kit": "^0.31.1",
|
||||
"typescript": "^5"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,67 @@
|
||||
import { Hono } from 'hono'
|
||||
import { cors } from 'hono/cors'
|
||||
import { secureHeaders } from 'hono/secure-headers'
|
||||
import { mcpAuthRouter } from '@hono/mcp'
|
||||
import type { ServerConfig } from './config'
|
||||
import type { AppBindings } from './context'
|
||||
import type { Db } from './db'
|
||||
import { createOAuthServerProvider } from './security/oauth-provider'
|
||||
import { authenticateRequest } from './security/http-auth'
|
||||
import { createAdminRoutes } from './routes/admin'
|
||||
import { createAuthRoutes } from './routes/auth'
|
||||
import { createMcpRoutes } from './routes/mcp'
|
||||
import { createOAuthRoutes } from './routes/oauth'
|
||||
import { createRpcWebSocket } from './rpc/ws'
|
||||
|
||||
export const createServerApp = (db: Db, config: ServerConfig) => {
|
||||
const app = new Hono<AppBindings>()
|
||||
const rpcWebSocket = createRpcWebSocket(db, config)
|
||||
const oauthProvider = createOAuthServerProvider(db, config)
|
||||
|
||||
app.use('*', async (c, next) => {
|
||||
const auth = await authenticateRequest(db, config, c.req.raw)
|
||||
c.set('db', db)
|
||||
c.set('config', config)
|
||||
c.set('auth', auth)
|
||||
c.set('authenticated', auth.authenticated)
|
||||
await next()
|
||||
})
|
||||
|
||||
app.use(
|
||||
'*',
|
||||
secureHeaders({
|
||||
crossOriginEmbedderPolicy: false,
|
||||
}),
|
||||
)
|
||||
app.use(
|
||||
'*',
|
||||
cors({
|
||||
origin: config.corsOrigin,
|
||||
credentials: true,
|
||||
}),
|
||||
)
|
||||
|
||||
app.get('/health', (c) => c.json({ ok: true }))
|
||||
app.route(
|
||||
'/',
|
||||
mcpAuthRouter({
|
||||
issuerUrl: new URL(config.origin),
|
||||
baseUrl: new URL(config.origin),
|
||||
resourceServerUrl: new URL('/api/mcp', config.origin),
|
||||
resourceName: 'CSBIE MCP',
|
||||
scopesSupported: ['mcp'],
|
||||
provider: oauthProvider,
|
||||
authorizationOptions: { rateLimit: false },
|
||||
tokenOptions: { rateLimit: false },
|
||||
clientRegistrationOptions: { rateLimit: false },
|
||||
revocationOptions: { rateLimit: false },
|
||||
}),
|
||||
)
|
||||
app.get('/ws', rpcWebSocket.upgradeWebSocket)
|
||||
app.route('/auth', createAuthRoutes())
|
||||
app.route('/admin', createAdminRoutes())
|
||||
app.route('/oauth', createOAuthRoutes())
|
||||
app.route('/mcp', createMcpRoutes())
|
||||
|
||||
return { app, websocket: rpcWebSocket.websocket }
|
||||
}
|
||||
@@ -0,0 +1,42 @@
|
||||
import { mkdirSync } from 'node:fs'
|
||||
import { dirname, resolve } from 'node:path'
|
||||
|
||||
export type ServerConfig = {
|
||||
port: number
|
||||
databasePath: string
|
||||
corsOrigin: string
|
||||
sessionCookieName: string
|
||||
rpName: string
|
||||
rpId: string
|
||||
origin: string
|
||||
authBaseUrl?: string
|
||||
mtsBaseUrl?: string
|
||||
izanagiBaseUrl?: string
|
||||
}
|
||||
|
||||
const optionalUrl = (value: string | undefined) => {
|
||||
if (!value) return undefined
|
||||
return new URL(value).toString()
|
||||
}
|
||||
|
||||
export const loadConfig = (): ServerConfig => {
|
||||
const port = Number(process.env.PORT ?? process.env.CSBIE_SERVER_PORT ?? 8787)
|
||||
const databasePath = resolve(process.env.CSBIE_DATABASE_PATH ?? './data/csbie.sqlite')
|
||||
mkdirSync(dirname(databasePath), { recursive: true })
|
||||
|
||||
const origin = process.env.CSBIE_ORIGIN ?? `http://localhost:${port}`
|
||||
const rpId = process.env.CSBIE_RP_ID ?? new URL(origin).hostname
|
||||
|
||||
return {
|
||||
port,
|
||||
databasePath,
|
||||
corsOrigin: process.env.CSBIE_CORS_ORIGIN ?? origin,
|
||||
sessionCookieName: process.env.CSBIE_SESSION_COOKIE ?? 'csbie_session',
|
||||
rpName: process.env.CSBIE_RP_NAME ?? 'CSBIE',
|
||||
rpId,
|
||||
origin,
|
||||
authBaseUrl: optionalUrl(process.env.SBI_AUTH_BASE_URL),
|
||||
mtsBaseUrl: optionalUrl(process.env.SBI_MTS_BASE_URL),
|
||||
izanagiBaseUrl: optionalUrl(process.env.SBI_IZANAGI_BASE_URL),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
import type { ServerConfig } from './config'
|
||||
import type { Db } from './db'
|
||||
|
||||
export type AppBindings = {
|
||||
Variables: {
|
||||
db: Db
|
||||
config: ServerConfig
|
||||
authenticated: boolean
|
||||
auth: AuthContext
|
||||
}
|
||||
}
|
||||
|
||||
export type AuthContext =
|
||||
| {
|
||||
type: 'none'
|
||||
authenticated: false
|
||||
}
|
||||
| {
|
||||
type: 'session'
|
||||
authenticated: true
|
||||
sessionId: string
|
||||
}
|
||||
| {
|
||||
type: 'apiKey'
|
||||
authenticated: true
|
||||
apiKeyId: string
|
||||
}
|
||||
@@ -0,0 +1,152 @@
|
||||
import { Database } from 'bun:sqlite'
|
||||
import { drizzle } from 'drizzle-orm/bun-sqlite'
|
||||
import * as schema from './schema'
|
||||
|
||||
export type Db = ReturnType<typeof createDb>
|
||||
|
||||
export const createDb = (path: string) => {
|
||||
const sqlite = new Database(path, { create: true, strict: true })
|
||||
sqlite.run('PRAGMA journal_mode = WAL')
|
||||
sqlite.run('PRAGMA foreign_keys = ON')
|
||||
sqlite.run(`
|
||||
CREATE TABLE IF NOT EXISTS app_state (
|
||||
key TEXT PRIMARY KEY,
|
||||
value TEXT NOT NULL,
|
||||
updated_at INTEGER NOT NULL
|
||||
)
|
||||
`)
|
||||
sqlite.run(`
|
||||
CREATE TABLE IF NOT EXISTS user_passkeys (
|
||||
id TEXT PRIMARY KEY,
|
||||
credential_id TEXT NOT NULL,
|
||||
public_key TEXT NOT NULL,
|
||||
counter INTEGER NOT NULL DEFAULT 0,
|
||||
transports TEXT,
|
||||
created_at INTEGER NOT NULL,
|
||||
updated_at INTEGER NOT NULL
|
||||
)
|
||||
`)
|
||||
sqlite.run(`
|
||||
CREATE TABLE IF NOT EXISTS passkey_challenges (
|
||||
id TEXT PRIMARY KEY,
|
||||
kind TEXT NOT NULL,
|
||||
challenge TEXT NOT NULL,
|
||||
expires_at INTEGER NOT NULL,
|
||||
created_at INTEGER NOT NULL
|
||||
)
|
||||
`)
|
||||
sqlite.run(`
|
||||
CREATE TABLE IF NOT EXISTS sessions (
|
||||
id TEXT PRIMARY KEY,
|
||||
expires_at INTEGER NOT NULL,
|
||||
created_at INTEGER NOT NULL
|
||||
)
|
||||
`)
|
||||
sqlite.run(`
|
||||
CREATE TABLE IF NOT EXISTS sbi_passkeys (
|
||||
id TEXT PRIMARY KEY,
|
||||
label TEXT NOT NULL,
|
||||
keyring_account TEXT NOT NULL,
|
||||
created_at INTEGER NOT NULL,
|
||||
updated_at INTEGER NOT NULL
|
||||
)
|
||||
`)
|
||||
sqlite.run(`
|
||||
CREATE TABLE IF NOT EXISTS api_keys (
|
||||
id TEXT PRIMARY KEY,
|
||||
label TEXT NOT NULL,
|
||||
token_hash TEXT NOT NULL UNIQUE,
|
||||
max_trades_per_hour INTEGER,
|
||||
max_trades_per_6_hours INTEGER,
|
||||
max_trades_per_day INTEGER,
|
||||
max_order_price_jpy INTEGER,
|
||||
max_order_amount_jpy INTEGER,
|
||||
allowed_methods TEXT,
|
||||
created_at INTEGER NOT NULL,
|
||||
last_used_at INTEGER,
|
||||
revoked_at INTEGER
|
||||
)
|
||||
`)
|
||||
const apiKeyColumns = sqlite
|
||||
.query<{ name: string }, []>('PRAGMA table_info(api_keys)')
|
||||
.all()
|
||||
.map((column) => column.name)
|
||||
for (const [name, type] of [
|
||||
['max_trades_per_hour', 'INTEGER'],
|
||||
['max_trades_per_6_hours', 'INTEGER'],
|
||||
['max_trades_per_day', 'INTEGER'],
|
||||
['max_order_price_jpy', 'INTEGER'],
|
||||
['max_order_amount_jpy', 'INTEGER'],
|
||||
['allowed_methods', 'TEXT'],
|
||||
] as const) {
|
||||
if (!apiKeyColumns.includes(name)) sqlite.run(`ALTER TABLE api_keys ADD COLUMN ${name} ${type}`)
|
||||
}
|
||||
sqlite.run(`
|
||||
CREATE TABLE IF NOT EXISTS api_key_trade_usage (
|
||||
api_key_id TEXT NOT NULL,
|
||||
window TEXT NOT NULL,
|
||||
hour_bucket TEXT NOT NULL,
|
||||
trade_count INTEGER NOT NULL DEFAULT 0,
|
||||
updated_at INTEGER NOT NULL,
|
||||
PRIMARY KEY (api_key_id, window, hour_bucket)
|
||||
)
|
||||
`)
|
||||
const tradeUsageColumns = sqlite
|
||||
.query<{ name: string }, []>('PRAGMA table_info(api_key_trade_usage)')
|
||||
.all()
|
||||
.map((column) => column.name)
|
||||
if (!tradeUsageColumns.includes('window')) {
|
||||
sqlite.run('DROP TABLE api_key_trade_usage')
|
||||
sqlite.run(`
|
||||
CREATE TABLE api_key_trade_usage (
|
||||
api_key_id TEXT NOT NULL,
|
||||
window TEXT NOT NULL,
|
||||
hour_bucket TEXT NOT NULL,
|
||||
trade_count INTEGER NOT NULL DEFAULT 0,
|
||||
updated_at INTEGER NOT NULL,
|
||||
PRIMARY KEY (api_key_id, window, hour_bucket)
|
||||
)
|
||||
`)
|
||||
}
|
||||
sqlite.run(`
|
||||
CREATE TABLE IF NOT EXISTS oauth_clients (
|
||||
id TEXT PRIMARY KEY,
|
||||
client TEXT NOT NULL,
|
||||
created_at INTEGER NOT NULL
|
||||
)
|
||||
`)
|
||||
sqlite.run(`
|
||||
CREATE TABLE IF NOT EXISTS oauth_authorization_codes (
|
||||
code TEXT PRIMARY KEY,
|
||||
client_id TEXT NOT NULL,
|
||||
redirect_uri TEXT NOT NULL,
|
||||
code_challenge TEXT NOT NULL,
|
||||
scopes TEXT NOT NULL,
|
||||
resource TEXT,
|
||||
api_key_settings TEXT,
|
||||
expires_at INTEGER NOT NULL,
|
||||
created_at INTEGER NOT NULL
|
||||
)
|
||||
`)
|
||||
const oauthCodeColumns = sqlite
|
||||
.query<{ name: string }, []>('PRAGMA table_info(oauth_authorization_codes)')
|
||||
.all()
|
||||
.map((column) => column.name)
|
||||
if (!oauthCodeColumns.includes('api_key_settings')) {
|
||||
sqlite.run('ALTER TABLE oauth_authorization_codes ADD COLUMN api_key_settings TEXT')
|
||||
}
|
||||
sqlite.run(`
|
||||
CREATE TABLE IF NOT EXISTS oauth_refresh_tokens (
|
||||
token_hash TEXT PRIMARY KEY,
|
||||
client_id TEXT NOT NULL,
|
||||
api_key_id TEXT NOT NULL,
|
||||
scopes TEXT NOT NULL,
|
||||
resource TEXT,
|
||||
expires_at INTEGER NOT NULL,
|
||||
revoked_at INTEGER,
|
||||
created_at INTEGER NOT NULL
|
||||
)
|
||||
`)
|
||||
|
||||
return drizzle(sqlite, { schema })
|
||||
}
|
||||
@@ -0,0 +1,102 @@
|
||||
import { integer, primaryKey, sqliteTable, text, uniqueIndex } from 'drizzle-orm/sqlite-core'
|
||||
|
||||
export const appState = sqliteTable('app_state', {
|
||||
key: text('key').primaryKey(),
|
||||
value: text('value').notNull(),
|
||||
updatedAt: integer('updated_at', { mode: 'timestamp_ms' }).notNull(),
|
||||
})
|
||||
|
||||
export const userPasskeys = sqliteTable('user_passkeys', {
|
||||
id: text('id').primaryKey(),
|
||||
credentialId: text('credential_id').notNull(),
|
||||
publicKey: text('public_key').notNull(),
|
||||
counter: integer('counter').notNull().default(0),
|
||||
transports: text('transports', { mode: 'json' }).$type<string[] | undefined>(),
|
||||
createdAt: integer('created_at', { mode: 'timestamp_ms' }).notNull(),
|
||||
updatedAt: integer('updated_at', { mode: 'timestamp_ms' }).notNull(),
|
||||
})
|
||||
|
||||
export const passkeyChallenges = sqliteTable('passkey_challenges', {
|
||||
id: text('id').primaryKey(),
|
||||
kind: text('kind', { enum: ['registration', 'authentication'] }).notNull(),
|
||||
challenge: text('challenge').notNull(),
|
||||
expiresAt: integer('expires_at', { mode: 'timestamp_ms' }).notNull(),
|
||||
createdAt: integer('created_at', { mode: 'timestamp_ms' }).notNull(),
|
||||
})
|
||||
|
||||
export const sessions = sqliteTable('sessions', {
|
||||
id: text('id').primaryKey(),
|
||||
expiresAt: integer('expires_at', { mode: 'timestamp_ms' }).notNull(),
|
||||
createdAt: integer('created_at', { mode: 'timestamp_ms' }).notNull(),
|
||||
})
|
||||
|
||||
export const sbiPasskeys = sqliteTable('sbi_passkeys', {
|
||||
id: text('id').primaryKey(),
|
||||
label: text('label').notNull(),
|
||||
keyringAccount: text('keyring_account').notNull(),
|
||||
createdAt: integer('created_at', { mode: 'timestamp_ms' }).notNull(),
|
||||
updatedAt: integer('updated_at', { mode: 'timestamp_ms' }).notNull(),
|
||||
})
|
||||
|
||||
export const apiKeys = sqliteTable(
|
||||
'api_keys',
|
||||
{
|
||||
id: text('id').primaryKey(),
|
||||
label: text('label').notNull(),
|
||||
tokenHash: text('token_hash').notNull(),
|
||||
maxTradesPerHour: integer('max_trades_per_hour'),
|
||||
maxTradesPer6Hours: integer('max_trades_per_6_hours'),
|
||||
maxTradesPerDay: integer('max_trades_per_day'),
|
||||
maxOrderPriceJpy: integer('max_order_price_jpy'),
|
||||
maxOrderAmountJpy: integer('max_order_amount_jpy'),
|
||||
allowedMethods: text('allowed_methods', { mode: 'json' }).$type<string[] | null>(),
|
||||
createdAt: integer('created_at', { mode: 'timestamp_ms' }).notNull(),
|
||||
lastUsedAt: integer('last_used_at', { mode: 'timestamp_ms' }),
|
||||
revokedAt: integer('revoked_at', { mode: 'timestamp_ms' }),
|
||||
},
|
||||
(table) => [uniqueIndex('api_keys_token_hash_unique').on(table.tokenHash)],
|
||||
)
|
||||
|
||||
export const apiKeyTradeUsage = sqliteTable(
|
||||
'api_key_trade_usage',
|
||||
{
|
||||
apiKeyId: text('api_key_id').notNull(),
|
||||
window: text('window', { enum: ['1h', '3h', '1d'] }).notNull(),
|
||||
hourBucket: text('hour_bucket').notNull(),
|
||||
tradeCount: integer('trade_count').notNull().default(0),
|
||||
updatedAt: integer('updated_at', { mode: 'timestamp_ms' }).notNull(),
|
||||
},
|
||||
(table) => [primaryKey({ columns: [table.apiKeyId, table.window, table.hourBucket] })],
|
||||
)
|
||||
|
||||
export const oauthClients = sqliteTable('oauth_clients', {
|
||||
id: text('id').primaryKey(),
|
||||
client: text('client', { mode: 'json' }).$type<Record<string, unknown>>().notNull(),
|
||||
createdAt: integer('created_at', { mode: 'timestamp_ms' }).notNull(),
|
||||
})
|
||||
|
||||
export const oauthAuthorizationCodes = sqliteTable('oauth_authorization_codes', {
|
||||
code: text('code').primaryKey(),
|
||||
clientId: text('client_id').notNull(),
|
||||
redirectUri: text('redirect_uri').notNull(),
|
||||
codeChallenge: text('code_challenge').notNull(),
|
||||
scopes: text('scopes', { mode: 'json' }).$type<string[]>().notNull(),
|
||||
resource: text('resource'),
|
||||
apiKeySettings: text('api_key_settings', { mode: 'json' }).$type<Record<
|
||||
string,
|
||||
unknown
|
||||
> | null>(),
|
||||
expiresAt: integer('expires_at', { mode: 'timestamp_ms' }).notNull(),
|
||||
createdAt: integer('created_at', { mode: 'timestamp_ms' }).notNull(),
|
||||
})
|
||||
|
||||
export const oauthRefreshTokens = sqliteTable('oauth_refresh_tokens', {
|
||||
tokenHash: text('token_hash').primaryKey(),
|
||||
clientId: text('client_id').notNull(),
|
||||
apiKeyId: text('api_key_id').notNull(),
|
||||
scopes: text('scopes', { mode: 'json' }).$type<string[]>().notNull(),
|
||||
resource: text('resource'),
|
||||
expiresAt: integer('expires_at', { mode: 'timestamp_ms' }).notNull(),
|
||||
revokedAt: integer('revoked_at', { mode: 'timestamp_ms' }),
|
||||
createdAt: integer('created_at', { mode: 'timestamp_ms' }).notNull(),
|
||||
})
|
||||
@@ -0,0 +1,17 @@
|
||||
import { loadConfig } from './config'
|
||||
import { createDb } from './db'
|
||||
import { createServerApp } from './app'
|
||||
|
||||
const config = loadConfig()
|
||||
const db = createDb(config.databasePath)
|
||||
const { app, websocket } = createServerApp(db, config)
|
||||
|
||||
const server = Bun.serve({
|
||||
port: config.port,
|
||||
fetch(request, server) {
|
||||
return app.fetch(request, { server })
|
||||
},
|
||||
websocket,
|
||||
})
|
||||
|
||||
console.log(`csbie-server listening on http://localhost:${server.port}`)
|
||||
@@ -0,0 +1,108 @@
|
||||
import { eq } from 'drizzle-orm'
|
||||
import { Hono } from 'hono'
|
||||
import type { MiddlewareHandler } from 'hono'
|
||||
import type { PlaintextStoredWebAuthnCredential } from '@repo/sbi-client'
|
||||
import type { AppBindings } from '../context'
|
||||
import { sbiPasskeys } from '../db/schema'
|
||||
import {
|
||||
createApiKey,
|
||||
listApiKeys,
|
||||
revokeApiKey,
|
||||
updateApiKeySettings,
|
||||
type ApiKeySettings,
|
||||
} from '../security/api-keys'
|
||||
import { randomId } from '../security/crypto'
|
||||
import { deleteSecret, saveSecret } from '../security/keyring'
|
||||
|
||||
export type StoredSbiPasskeySecret = {
|
||||
credential: PlaintextStoredWebAuthnCredential
|
||||
tradePassword?: string
|
||||
deviceId?: string
|
||||
}
|
||||
|
||||
const requireAuth: MiddlewareHandler<AppBindings> = async (c, next) => {
|
||||
if (!c.get('authenticated')) return c.json({ error: 'unauthorized' }, 401)
|
||||
await next()
|
||||
}
|
||||
|
||||
export const createAdminRoutes = () => {
|
||||
const app = new Hono<AppBindings>()
|
||||
app.use('*', requireAuth)
|
||||
|
||||
app.get('/api-keys', async (c) => c.json({ apiKeys: await listApiKeys(c.get('db')) }))
|
||||
|
||||
app.post('/api-keys', async (c) => {
|
||||
const { label, settings } = await c.req.json<{ label?: string; settings?: ApiKeySettings }>()
|
||||
if (!label?.trim()) return c.json({ error: 'label is required' }, 400)
|
||||
const key = await createApiKey(c.get('db'), label.trim(), settings)
|
||||
return c.json({ apiKey: key }, 201)
|
||||
})
|
||||
|
||||
app.patch('/api-keys/:id/settings', async (c) => {
|
||||
const body = await c.req.json<ApiKeySettings>()
|
||||
await updateApiKeySettings(c.get('db'), c.req.param('id'), body)
|
||||
return c.json({ ok: true })
|
||||
})
|
||||
|
||||
app.delete('/api-keys/:id', async (c) => {
|
||||
await revokeApiKey(c.get('db'), c.req.param('id'))
|
||||
return c.json({ ok: true })
|
||||
})
|
||||
|
||||
app.get('/sbi-passkeys', async (c) => {
|
||||
const rows = await c.get('db').select().from(sbiPasskeys).orderBy(sbiPasskeys.createdAt)
|
||||
return c.json({
|
||||
passkeys: rows.map(({ keyringAccount: _keyringAccount, ...row }) => ({
|
||||
...row,
|
||||
keyringAccount: undefined,
|
||||
})),
|
||||
})
|
||||
})
|
||||
|
||||
app.post('/sbi-passkeys', async (c) => {
|
||||
const body = await c.req.json<{
|
||||
label?: string
|
||||
credential?: PlaintextStoredWebAuthnCredential
|
||||
tradePassword?: string
|
||||
deviceId?: string
|
||||
}>()
|
||||
if (!body.label?.trim() || !body.credential) {
|
||||
return c.json({ error: 'label and credential are required' }, 400)
|
||||
}
|
||||
|
||||
const now = new Date()
|
||||
const id = randomId('sbi')
|
||||
const keyringAccount = `sbi-passkey:${id}`
|
||||
await saveSecret(keyringAccount, {
|
||||
credential: body.credential,
|
||||
tradePassword: body.tradePassword,
|
||||
deviceId: body.deviceId,
|
||||
} satisfies StoredSbiPasskeySecret)
|
||||
await c.get('db').insert(sbiPasskeys).values({
|
||||
id,
|
||||
label: body.label.trim(),
|
||||
keyringAccount,
|
||||
createdAt: now,
|
||||
updatedAt: now,
|
||||
})
|
||||
|
||||
return c.json(
|
||||
{ passkey: { id, label: body.label.trim(), createdAt: now, updatedAt: now } },
|
||||
201,
|
||||
)
|
||||
})
|
||||
|
||||
app.delete('/sbi-passkeys/:id', async (c) => {
|
||||
const db = c.get('db')
|
||||
const [row] = await db
|
||||
.select()
|
||||
.from(sbiPasskeys)
|
||||
.where(eq(sbiPasskeys.id, c.req.param('id')))
|
||||
if (!row) return c.json({ error: 'not found' }, 404)
|
||||
await deleteSecret(row.keyringAccount)
|
||||
await db.delete(sbiPasskeys).where(eq(sbiPasskeys.id, row.id))
|
||||
return c.json({ ok: true })
|
||||
})
|
||||
|
||||
return app
|
||||
}
|
||||
@@ -0,0 +1,218 @@
|
||||
import {
|
||||
type AuthenticatorTransportFuture,
|
||||
generateAuthenticationOptions,
|
||||
generateRegistrationOptions,
|
||||
verifyAuthenticationResponse,
|
||||
verifyRegistrationResponse,
|
||||
} from '@simplewebauthn/server'
|
||||
import { eq } from 'drizzle-orm'
|
||||
import { Hono } from 'hono'
|
||||
import type { AppBindings } from '../context'
|
||||
import { appState, passkeyChallenges, userPasskeys } from '../db/schema'
|
||||
import { randomId } from '../security/crypto'
|
||||
import {
|
||||
clearSessionCookie,
|
||||
createSession,
|
||||
setSessionCookie,
|
||||
verifySessionCookie,
|
||||
} from '../security/sessions'
|
||||
import { verifySetupPassword } from '../security/setup'
|
||||
|
||||
type PasskeyRow = typeof userPasskeys.$inferSelect
|
||||
|
||||
const getConfigured = async (db: AppBindings['Variables']['db']) => {
|
||||
const [row] = await db.select().from(appState).where(eq(appState.key, 'configured')).limit(1)
|
||||
return row?.value === 'true'
|
||||
}
|
||||
|
||||
const saveChallenge = async (
|
||||
db: AppBindings['Variables']['db'],
|
||||
kind: 'registration' | 'authentication',
|
||||
challenge: string,
|
||||
) => {
|
||||
const now = new Date()
|
||||
const id = randomId('chal')
|
||||
await db.insert(passkeyChallenges).values({
|
||||
id,
|
||||
kind,
|
||||
challenge,
|
||||
createdAt: now,
|
||||
expiresAt: new Date(now.getTime() + 5 * 60 * 1000),
|
||||
})
|
||||
return id
|
||||
}
|
||||
|
||||
const consumeChallenge = async (
|
||||
db: AppBindings['Variables']['db'],
|
||||
id: string,
|
||||
kind: 'registration' | 'authentication',
|
||||
) => {
|
||||
const [row] = await db
|
||||
.select()
|
||||
.from(passkeyChallenges)
|
||||
.where(eq(passkeyChallenges.id, id))
|
||||
.limit(1)
|
||||
if (!row || row.kind !== kind || row.expiresAt < new Date()) throw new Error('challenge expired')
|
||||
await db.delete(passkeyChallenges).where(eq(passkeyChallenges.id, id))
|
||||
return row.challenge
|
||||
}
|
||||
|
||||
const toCredentialDescriptor = (row: PasskeyRow) => ({
|
||||
id: row.credentialId,
|
||||
transports: row.transports?.filter(isAuthenticatorTransport),
|
||||
})
|
||||
|
||||
const isAuthenticatorTransport = (value: string): value is AuthenticatorTransportFuture =>
|
||||
['ble', 'cable', 'hybrid', 'internal', 'nfc', 'smart-card', 'usb'].includes(value)
|
||||
|
||||
const passkeyForCredential = async (db: AppBindings['Variables']['db'], credentialId: string) => {
|
||||
const [row] = await db
|
||||
.select()
|
||||
.from(userPasskeys)
|
||||
.where(eq(userPasskeys.credentialId, credentialId))
|
||||
.limit(1)
|
||||
return row
|
||||
}
|
||||
|
||||
export const createAuthRoutes = () => {
|
||||
const app = new Hono<AppBindings>()
|
||||
|
||||
app.get('/status', async (c) => {
|
||||
const db = c.get('db')
|
||||
return c.json({
|
||||
configured: await getConfigured(db),
|
||||
authenticated: await verifySessionCookie(c, db, c.get('config')),
|
||||
})
|
||||
})
|
||||
|
||||
app.post('/setup/options', async (c) => {
|
||||
const db = c.get('db')
|
||||
if (await getConfigured(db)) return c.json({ error: 'already configured' }, 409)
|
||||
|
||||
const { password } = await c.req.json<{ password?: string }>()
|
||||
if (!password || !verifySetupPassword(password)) {
|
||||
return c.json({ error: 'invalid setup password' }, 401)
|
||||
}
|
||||
|
||||
const config = c.get('config')
|
||||
const options = await generateRegistrationOptions({
|
||||
rpName: config.rpName,
|
||||
rpID: config.rpId,
|
||||
userName: 'owner',
|
||||
userDisplayName: 'Owner',
|
||||
attestationType: 'none',
|
||||
authenticatorSelection: {
|
||||
residentKey: 'preferred',
|
||||
userVerification: 'required',
|
||||
},
|
||||
})
|
||||
const challengeId = await saveChallenge(db, 'registration', options.challenge)
|
||||
return c.json({ options, challengeId })
|
||||
})
|
||||
|
||||
app.post('/setup/verify', async (c) => {
|
||||
const db = c.get('db')
|
||||
if (await getConfigured(db)) return c.json({ error: 'already configured' }, 409)
|
||||
|
||||
const { challengeId, response } = await c.req.json<{
|
||||
challengeId?: string
|
||||
response?: unknown
|
||||
}>()
|
||||
if (!challengeId || !response) return c.json({ error: 'missing registration response' }, 400)
|
||||
|
||||
const expectedChallenge = await consumeChallenge(db, challengeId, 'registration')
|
||||
const config = c.get('config')
|
||||
const verification = await verifyRegistrationResponse({
|
||||
response: response as never,
|
||||
expectedChallenge,
|
||||
expectedOrigin: config.origin,
|
||||
expectedRPID: config.rpId,
|
||||
requireUserVerification: true,
|
||||
})
|
||||
|
||||
if (!verification.verified || !verification.registrationInfo) {
|
||||
return c.json({ error: 'registration failed' }, 400)
|
||||
}
|
||||
|
||||
const now = new Date()
|
||||
const credential = verification.registrationInfo.credential
|
||||
await db.insert(userPasskeys).values({
|
||||
id: randomId('upk'),
|
||||
credentialId: credential.id,
|
||||
publicKey: Buffer.from(credential.publicKey).toString('base64url'),
|
||||
counter: credential.counter,
|
||||
transports: (response as { response?: { transports?: string[] } }).response?.transports,
|
||||
createdAt: now,
|
||||
updatedAt: now,
|
||||
})
|
||||
await db
|
||||
.insert(appState)
|
||||
.values({ key: 'configured', value: 'true', updatedAt: now })
|
||||
.onConflictDoUpdate({ target: appState.key, set: { value: 'true', updatedAt: now } })
|
||||
|
||||
const session = await createSession(db)
|
||||
setSessionCookie(c, config, session.id, session.expiresAt)
|
||||
return c.json({ ok: true })
|
||||
})
|
||||
|
||||
app.post('/login/options', async (c) => {
|
||||
const db = c.get('db')
|
||||
const rows = await db.select().from(userPasskeys)
|
||||
if (rows.length === 0) return c.json({ error: 'not configured' }, 409)
|
||||
|
||||
const options = await generateAuthenticationOptions({
|
||||
rpID: c.get('config').rpId,
|
||||
allowCredentials: rows.map(toCredentialDescriptor),
|
||||
userVerification: 'required',
|
||||
})
|
||||
const challengeId = await saveChallenge(db, 'authentication', options.challenge)
|
||||
return c.json({ options, challengeId })
|
||||
})
|
||||
|
||||
app.post('/login/verify', async (c) => {
|
||||
const db = c.get('db')
|
||||
const { challengeId, response } = await c.req.json<{
|
||||
challengeId?: string
|
||||
response?: { id?: string }
|
||||
}>()
|
||||
if (!challengeId || !response?.id)
|
||||
return c.json({ error: 'missing authentication response' }, 400)
|
||||
|
||||
const passkey = await passkeyForCredential(db, response.id)
|
||||
if (!passkey) return c.json({ error: 'unknown passkey' }, 401)
|
||||
|
||||
const expectedChallenge = await consumeChallenge(db, challengeId, 'authentication')
|
||||
const config = c.get('config')
|
||||
const verification = await verifyAuthenticationResponse({
|
||||
response: response as never,
|
||||
expectedChallenge,
|
||||
expectedOrigin: config.origin,
|
||||
expectedRPID: config.rpId,
|
||||
credential: {
|
||||
id: passkey.credentialId,
|
||||
publicKey: Buffer.from(passkey.publicKey, 'base64url'),
|
||||
counter: passkey.counter,
|
||||
transports: passkey.transports?.filter(isAuthenticatorTransport),
|
||||
},
|
||||
requireUserVerification: true,
|
||||
})
|
||||
|
||||
if (!verification.verified) return c.json({ error: 'authentication failed' }, 401)
|
||||
|
||||
await db
|
||||
.update(userPasskeys)
|
||||
.set({ counter: verification.authenticationInfo.newCounter, updatedAt: new Date() })
|
||||
.where(eq(userPasskeys.id, passkey.id))
|
||||
|
||||
const session = await createSession(db)
|
||||
setSessionCookie(c, config, session.id, session.expiresAt)
|
||||
return c.json({ ok: true })
|
||||
})
|
||||
|
||||
app.post('/logout', async (c) => {
|
||||
await clearSessionCookie(c, c.get('db'), c.get('config'))
|
||||
return c.json({ ok: true })
|
||||
})
|
||||
|
||||
return app
|
||||
}
|
||||
@@ -0,0 +1,604 @@
|
||||
import { Hono } from 'hono'
|
||||
import type { Context } from 'hono'
|
||||
import { randomUUID } from 'node:crypto'
|
||||
import { StreamableHTTPTransport } from '@hono/mcp'
|
||||
import { McpServer } from '@modelcontextprotocol/sdk/server/mcp.js'
|
||||
import { eq } from 'drizzle-orm'
|
||||
import * as z from 'zod/v4'
|
||||
import type { AppBindings, AuthContext } from '../context'
|
||||
import { sbiPasskeys } from '../db/schema'
|
||||
import {
|
||||
invokeSbiMethod,
|
||||
isCashOrderMethod,
|
||||
isTradingMethod,
|
||||
RPC_METHODS,
|
||||
type RpcMethod,
|
||||
} from '../rpc/methods'
|
||||
import { connectSbi } from '../rpc/sbi-session'
|
||||
import type { StoredSbiPasskeySecret } from './admin'
|
||||
import {
|
||||
assertAndConsumeApiKeyTradeLimits,
|
||||
assertApiKeyMethodAllowed,
|
||||
} from '../security/trade-limits'
|
||||
import { readSecret } from '../security/keyring'
|
||||
import { effectiveSbiDeviceId, effectiveSbiTradePassword } from '../security/sbi-credentials'
|
||||
|
||||
const jsonText = (value: unknown) => JSON.stringify(value, null, 2)
|
||||
|
||||
const textResult = (value: unknown) => ({
|
||||
content: [{ type: 'text' as const, text: typeof value === 'string' ? value : jsonText(value) }],
|
||||
})
|
||||
|
||||
const requireAuthenticated = (auth: AuthContext) => {
|
||||
if (!auth.authenticated) throw new Error('unauthorized')
|
||||
}
|
||||
|
||||
const toolNameForMethod = (method: RpcMethod) => `csbie_sbi_${method.replaceAll('.', '_')}`
|
||||
|
||||
const ORDER_SUBMIT_TICKET_TTL_MS = 10 * 60 * 1000
|
||||
|
||||
type OrderSubmitTicket = {
|
||||
passkeyId: string
|
||||
estimateMethod: RpcMethod
|
||||
submitMethod: RpcMethod
|
||||
params: unknown
|
||||
confirmationId?: string
|
||||
authKey: string
|
||||
expiresAt: Date
|
||||
}
|
||||
|
||||
const orderSubmitTickets = new Map<string, OrderSubmitTicket>()
|
||||
|
||||
const authKey = (auth: AuthContext) => {
|
||||
if (auth.type === 'apiKey') return `apiKey:${auth.apiKeyId}`
|
||||
if (auth.type === 'session') return `session:${auth.sessionId}`
|
||||
return 'none'
|
||||
}
|
||||
|
||||
const cleanupExpiredOrderSubmitTickets = (now = new Date()) => {
|
||||
for (const [uuid, ticket] of orderSubmitTickets) {
|
||||
if (ticket.expiresAt <= now) orderSubmitTickets.delete(uuid)
|
||||
}
|
||||
}
|
||||
|
||||
const orderSubmitMethodByEstimateMethod = {
|
||||
'orders.cash.estimate': 'orders.cash.place',
|
||||
'orders.cash.estimateCorrection': 'orders.cash.placeCorrection',
|
||||
'orders.cash.estimateCorrectionConfirm': 'orders.cash.placeCorrection',
|
||||
'orders.cash.estimateCancel': 'orders.cash.placeCancel',
|
||||
'orders.margin.estimateOpen': 'orders.margin.open',
|
||||
'orders.margin.estimateClose': 'orders.margin.close',
|
||||
'orders.margin.estimateCloseSummary': 'orders.margin.closeSummary',
|
||||
'orders.margin.estimateSummary': 'orders.margin.placeSummary',
|
||||
'orders.margin.estimateActualDelivery': 'orders.margin.actualDelivery',
|
||||
'orders.ifd.estimate': 'orders.ifd.place',
|
||||
'orders.ifd.estimateCorrection': 'orders.ifd.placeCorrection',
|
||||
'orders.ifd.estimateCancel': 'orders.ifd.placeCancel',
|
||||
'orders.themeInvestment.estimate': 'orders.themeInvestment.place',
|
||||
} as const satisfies Partial<Record<RpcMethod, RpcMethod>>
|
||||
|
||||
const submitMethodForEstimateMethod = (method: RpcMethod) =>
|
||||
orderSubmitMethodByEstimateMethod[method as keyof typeof orderSubmitMethodByEstimateMethod]
|
||||
|
||||
const isDirectOrderSubmitMethod = (method: RpcMethod) => isTradingMethod(method)
|
||||
|
||||
const mcpExposedRpcMethods = RPC_METHODS.filter((method) => !isDirectOrderSubmitMethod(method))
|
||||
|
||||
const orderSubmitParams = (value: unknown, confirmationId?: string) => {
|
||||
if (!value || typeof value !== 'object' || Array.isArray(value)) return { allowTrading: true }
|
||||
return {
|
||||
...value,
|
||||
...(confirmationId ? { confirmationId } : {}),
|
||||
allowTrading: true,
|
||||
}
|
||||
}
|
||||
|
||||
const confirmationIdFromPreview = (value: unknown) => {
|
||||
if (!value || typeof value !== 'object' || Array.isArray(value)) return undefined
|
||||
const confirmationId = (value as Record<string, unknown>).confirmationId
|
||||
return typeof confirmationId === 'string' && confirmationId ? confirmationId : undefined
|
||||
}
|
||||
|
||||
const accountTypeSchema = z.enum(['general', 'specific', 'nisa', 'juniorNisa', 'unknown'])
|
||||
const depositTypeSchema = z.enum(['general', 'specific', 'nisa', 'juniorNisa', 'unknown'])
|
||||
const tradeSideSchema = z.enum(['buy', 'sell'])
|
||||
const marketCodeSchema = z.string().min(1).describe('SBI market code')
|
||||
const issueCodeSchema = z.string().min(1).describe('Issue code')
|
||||
const orderIdSchema = z.string().min(1).describe('Order id')
|
||||
const positionIdSchema = z.string().min(1).describe('Position id')
|
||||
|
||||
const pagingSchema = {
|
||||
index: z.number().int().min(0).optional().describe('Start index for the result list'),
|
||||
limit: z.number().int().positive().optional().describe('Maximum number of items to fetch'),
|
||||
}
|
||||
|
||||
const issueOptionsSchema = z.object({
|
||||
issueCode: issueCodeSchema,
|
||||
market: marketCodeSchema.optional(),
|
||||
})
|
||||
|
||||
const issueChartOptionsSchema = issueOptionsSchema.extend({
|
||||
period: z.enum(['minute', 'day', 'week', 'month']).optional().describe('Chart period'),
|
||||
unit: z
|
||||
.number()
|
||||
.int()
|
||||
.positive()
|
||||
.optional()
|
||||
.describe('Candle unit. Minute charts accept 1, 5, 10, or 15; other periods use 1'),
|
||||
count: z
|
||||
.number()
|
||||
.int()
|
||||
.positive()
|
||||
.max(9999)
|
||||
.optional()
|
||||
.describe('Number of historical prices to request'),
|
||||
})
|
||||
|
||||
const issueSearchOptionsSchema = z.object({
|
||||
query: z.string().min(1).describe('Search text, such as an issue code, name, or keyword'),
|
||||
market: marketCodeSchema.optional().describe('Client-side market code filter'),
|
||||
limit: z.number().int().positive().optional().describe('Maximum number of returned issues'),
|
||||
})
|
||||
|
||||
const cashPositionOptionsSchema = z.object({
|
||||
...pagingSchema,
|
||||
issueCode: issueCodeSchema.optional(),
|
||||
market: marketCodeSchema.optional(),
|
||||
accountType: accountTypeSchema.optional(),
|
||||
})
|
||||
|
||||
const marginPositionOptionsSchema = z.object({
|
||||
...pagingSchema,
|
||||
issueCode: issueCodeSchema.optional(),
|
||||
market: marketCodeSchema.optional(),
|
||||
side: tradeSideSchema.optional(),
|
||||
accountType: accountTypeSchema.optional(),
|
||||
})
|
||||
|
||||
const orderInquiryOptionsSchema = z.object({
|
||||
...pagingSchema,
|
||||
from: z.string().optional().describe('Start date for the inquiry range'),
|
||||
to: z.string().optional().describe('End date for the inquiry range'),
|
||||
issueCode: issueCodeSchema.optional(),
|
||||
market: marketCodeSchema.optional(),
|
||||
status: z.enum(['open', 'executed', 'cancelled', 'expired', 'rejected', 'unknown']).optional(),
|
||||
})
|
||||
|
||||
const boardOptionsSchema = issueOptionsSchema.extend({
|
||||
accountType: accountTypeSchema.optional(),
|
||||
side: z
|
||||
.enum([
|
||||
'cashBuy',
|
||||
'cashSell',
|
||||
'marginOpen',
|
||||
'marginOpenBuy',
|
||||
'marginOpenSell',
|
||||
'marginClose',
|
||||
'marginCloseBuy',
|
||||
'marginCloseSell',
|
||||
])
|
||||
.optional(),
|
||||
})
|
||||
|
||||
const stockOrderBaseSchema = z.object({
|
||||
issueCode: issueCodeSchema,
|
||||
market: marketCodeSchema,
|
||||
side: tradeSideSchema,
|
||||
accountType: accountTypeSchema.optional(),
|
||||
quantity: z.number().positive().describe('Order quantity'),
|
||||
depositType: depositTypeSchema.optional(),
|
||||
})
|
||||
|
||||
const cashOrderPriceConditionSchema = z.enum([
|
||||
'limit',
|
||||
'limitAtOpen',
|
||||
'limitAtClose',
|
||||
'limitIoc',
|
||||
'market',
|
||||
'marketAtOpen',
|
||||
'marketAtClose',
|
||||
'marketIoc',
|
||||
'funari',
|
||||
])
|
||||
|
||||
const cashOrderSchema = stockOrderBaseSchema.extend({
|
||||
price: z.number().positive().optional().describe('Order price for price-based orders'),
|
||||
kind: z.enum(['market', 'limit', 'stop', 'oco', 'ifd', 'ifdo', 's', 'unknown']).optional(),
|
||||
priceCondition: cashOrderPriceConditionSchema
|
||||
.optional()
|
||||
.describe('APK/MTS execution condition for cash orders'),
|
||||
orderTerm: z.enum(['day', 'week', 'date']).optional().describe('Order validity term'),
|
||||
orderDate: z
|
||||
.string()
|
||||
.optional()
|
||||
.describe('Validity date used when orderTerm is date, in yyyyMMdd or yyyy-MM-dd format'),
|
||||
orderMethod: z.enum(['normal', 'stop', 'oco']).optional().describe('Special order method'),
|
||||
triggerZone: z.enum(['above', 'below']).optional().describe('Stop trigger direction'),
|
||||
triggerPrice: z.number().positive().optional().describe('Stop trigger price'),
|
||||
secondaryPriceCondition: cashOrderPriceConditionSchema
|
||||
.optional()
|
||||
.describe('Secondary execution condition for OCO orders'),
|
||||
secondaryPrice: z.number().positive().optional().describe('Secondary price for OCO orders'),
|
||||
sorLastMarket: marketCodeSchema
|
||||
.optional()
|
||||
.describe('Previous market code sent with SOR orders; defaults to login profile'),
|
||||
})
|
||||
|
||||
const placeCashOrderSchema = cashOrderSchema.extend({
|
||||
confirmationId: z
|
||||
.string()
|
||||
.optional()
|
||||
.describe('Confirmation ID returned by the confirmation step'),
|
||||
allowTrading: z.literal(true).optional().describe('Explicitly allows sending a live order'),
|
||||
})
|
||||
|
||||
const orderCorrectionSchema = z.object({
|
||||
orderId: orderIdSchema,
|
||||
quantity: z.number().positive().optional().describe('Corrected order quantity'),
|
||||
price: z.number().positive().optional().describe('Corrected order price'),
|
||||
})
|
||||
|
||||
const placeOrderCorrectionSchema = orderCorrectionSchema.extend({
|
||||
allowTrading: z
|
||||
.literal(true)
|
||||
.optional()
|
||||
.describe('Explicitly allows sending a live correction request'),
|
||||
})
|
||||
|
||||
const orderCancelSchema = z.object({
|
||||
orderNumber: z.string().min(1).describe('Order number shown in order inquiry'),
|
||||
orderId: orderIdSchema.optional().describe('Original order id shown in order inquiry'),
|
||||
tradeId: z.string().min(1).optional().describe('Original trade id code'),
|
||||
cancelType: z.string().min(1).optional().describe('Additional cancel flag'),
|
||||
})
|
||||
|
||||
const placeOrderCancelSchema = orderCancelSchema.extend({
|
||||
tradePassword: z.string().optional().describe('Trading password used by SBI'),
|
||||
allowTrading: z
|
||||
.literal(true)
|
||||
.optional()
|
||||
.describe('Explicitly allows sending a live cancellation request'),
|
||||
})
|
||||
|
||||
const marginOpenOrderSchema = cashOrderSchema
|
||||
|
||||
const placeMarginOpenOrderSchema = marginOpenOrderSchema.extend({
|
||||
confirmationId: z
|
||||
.string()
|
||||
.optional()
|
||||
.describe('Confirmation ID returned by the confirmation step'),
|
||||
allowTrading: z
|
||||
.literal(true)
|
||||
.optional()
|
||||
.describe('Explicitly allows sending a live margin open order'),
|
||||
})
|
||||
|
||||
const marginCloseOrderSchema = cashOrderSchema.extend({
|
||||
positionId: positionIdSchema.optional().describe('Position ID to close'),
|
||||
})
|
||||
|
||||
const placeMarginCloseOrderSchema = marginCloseOrderSchema.extend({
|
||||
allowTrading: z
|
||||
.literal(true)
|
||||
.optional()
|
||||
.describe('Explicitly allows sending a live margin close order'),
|
||||
})
|
||||
|
||||
const actualDeliveryOrderSchema = z.object({
|
||||
issueCode: issueCodeSchema,
|
||||
market: marketCodeSchema,
|
||||
accountType: accountTypeSchema.optional(),
|
||||
quantity: z.number().positive().describe('Order quantity'),
|
||||
depositType: depositTypeSchema.optional(),
|
||||
price: z.number().positive().optional().describe('Order price for price-based requests'),
|
||||
kind: z.enum(['genbiki', 'genwatashi']),
|
||||
positionId: positionIdSchema.optional().describe('Position ID to deliver'),
|
||||
})
|
||||
|
||||
const placeActualDeliveryOrderSchema = actualDeliveryOrderSchema.extend({
|
||||
confirmationId: z
|
||||
.string()
|
||||
.optional()
|
||||
.describe('Confirmation ID returned by the confirmation step'),
|
||||
allowTrading: z
|
||||
.literal(true)
|
||||
.optional()
|
||||
.describe('Explicitly allows sending a live actual-delivery order'),
|
||||
})
|
||||
|
||||
const ifdOrderSchema = cashOrderSchema.extend({
|
||||
tradeType: z
|
||||
.enum(['cash', 'marginOpen'])
|
||||
.optional()
|
||||
.describe('Product to use for the first IFD leg'),
|
||||
})
|
||||
|
||||
const placeIfdOrderSchema = ifdOrderSchema.extend({
|
||||
confirmationId: z
|
||||
.string()
|
||||
.optional()
|
||||
.describe('Confirmation ID returned by the confirmation step'),
|
||||
allowTrading: z.literal(true).optional().describe('Explicitly allows sending a live IFD order'),
|
||||
})
|
||||
|
||||
const themeInvestmentOrderSchema = z.object({
|
||||
themeId: z.string().min(1).describe('Theme ID for the theme investment order'),
|
||||
side: tradeSideSchema,
|
||||
amount: z.number().positive().optional().describe('Order amount for the theme investment order'),
|
||||
})
|
||||
|
||||
const placeThemeInvestmentOrderSchema = themeInvestmentOrderSchema.extend({
|
||||
allowTrading: z
|
||||
.literal(true)
|
||||
.optional()
|
||||
.describe('Explicitly allows sending a live theme investment order'),
|
||||
})
|
||||
|
||||
const methodParamSchemas = {
|
||||
'session.profile': undefined,
|
||||
'account.profile': undefined,
|
||||
'account.power.buyingPower': undefined,
|
||||
'account.power.collateralRatio': undefined,
|
||||
'account.positions.cash': cashPositionOptionsSchema.optional(),
|
||||
'account.positions.cashDetail': cashPositionOptionsSchema.optional(),
|
||||
'account.positions.cashForIssue': issueOptionsSchema,
|
||||
'account.positions.margin': marginPositionOptionsSchema.optional(),
|
||||
'account.positions.marginDetail': marginPositionOptionsSchema.optional(),
|
||||
'account.positions.marginForIssue': issueOptionsSchema,
|
||||
'account.positions.marginSummaryForIssue': issueOptionsSchema,
|
||||
'account.positions.marginDetailsForIssue': issueOptionsSchema,
|
||||
'account.positions.closeableMargin': marginPositionOptionsSchema,
|
||||
'account.positions.deliverableMargin': marginPositionOptionsSchema,
|
||||
'account.profitLoss.unrealized': undefined,
|
||||
'market.issue.search': issueSearchOptionsSchema,
|
||||
'market.issue.suggest': issueSearchOptionsSchema,
|
||||
'market.issue.allowedPrices': issueOptionsSchema,
|
||||
'market.issue.board': issueOptionsSchema,
|
||||
'market.issue.chart': issueChartOptionsSchema,
|
||||
'market.issue.openOrders': issueOptionsSchema,
|
||||
'market.issue.tradingInfo': boardOptionsSchema,
|
||||
'market.index.major': undefined,
|
||||
'market.overview': undefined,
|
||||
'market.ranking.market': undefined,
|
||||
'market.ranking.sector': undefined,
|
||||
'market.ranking.sbi': undefined,
|
||||
'news.list': undefined,
|
||||
'watchlist.list': undefined,
|
||||
'orders.inquiry.executionsToday': orderInquiryOptionsSchema.optional(),
|
||||
'orders.inquiry.open': orderInquiryOptionsSchema.optional(),
|
||||
'orders.cash.estimate': cashOrderSchema,
|
||||
'orders.cash.place': placeCashOrderSchema,
|
||||
'orders.cash.estimateCorrection': orderCorrectionSchema,
|
||||
'orders.cash.estimateCorrectionConfirm': orderCorrectionSchema,
|
||||
'orders.cash.placeCorrection': placeOrderCorrectionSchema,
|
||||
'orders.cash.estimateCancel': orderCancelSchema,
|
||||
'orders.cash.placeCancel': placeOrderCancelSchema,
|
||||
'orders.margin.estimateOpen': marginOpenOrderSchema,
|
||||
'orders.margin.open': placeMarginOpenOrderSchema,
|
||||
'orders.margin.estimateClose': marginCloseOrderSchema,
|
||||
'orders.margin.close': placeMarginCloseOrderSchema,
|
||||
'orders.margin.estimateCloseSummary': marginCloseOrderSchema,
|
||||
'orders.margin.closeSummary': placeMarginCloseOrderSchema,
|
||||
'orders.margin.estimateSummary': marginCloseOrderSchema,
|
||||
'orders.margin.placeSummary': placeMarginCloseOrderSchema,
|
||||
'orders.margin.estimateActualDelivery': actualDeliveryOrderSchema,
|
||||
'orders.margin.actualDelivery': placeActualDeliveryOrderSchema,
|
||||
'orders.ifd.estimate': ifdOrderSchema,
|
||||
'orders.ifd.place': placeIfdOrderSchema,
|
||||
'orders.ifd.estimateCorrection': orderCorrectionSchema,
|
||||
'orders.ifd.placeCorrection': placeOrderCorrectionSchema,
|
||||
'orders.ifd.estimateCancel': orderCorrectionSchema,
|
||||
'orders.ifd.placeCancel': placeOrderCorrectionSchema,
|
||||
'orders.themeInvestment.list': undefined,
|
||||
'orders.themeInvestment.estimate': themeInvestmentOrderSchema,
|
||||
'orders.themeInvestment.place': placeThemeInvestmentOrderSchema,
|
||||
} satisfies Record<RpcMethod, z.ZodType | undefined>
|
||||
|
||||
const createMcpServer = (c: Context<AppBindings>) => {
|
||||
const db = c.get('db')
|
||||
const config = c.get('config')
|
||||
const auth = c.get('auth')
|
||||
|
||||
const server = new McpServer({
|
||||
name: 'csbie',
|
||||
version: '0.1.0',
|
||||
})
|
||||
|
||||
server.registerTool(
|
||||
'csbie_sbi_methods',
|
||||
{
|
||||
title: 'List SBI RPC Methods',
|
||||
description: 'List SBI client methods exposed through CSBIE.',
|
||||
inputSchema: {},
|
||||
},
|
||||
async () => {
|
||||
requireAuthenticated(auth)
|
||||
return textResult({
|
||||
methods: mcpExposedRpcMethods,
|
||||
submitTool: 'csbie_sbi_submit_order',
|
||||
})
|
||||
},
|
||||
)
|
||||
|
||||
server.registerTool(
|
||||
'csbie_sbi_passkeys',
|
||||
{
|
||||
title: 'List SBI Passkeys',
|
||||
description: 'List saved SBI passkey profiles. Secret material is never returned.',
|
||||
inputSchema: {},
|
||||
},
|
||||
async () => {
|
||||
requireAuthenticated(auth)
|
||||
const rows = await db
|
||||
.select({
|
||||
id: sbiPasskeys.id,
|
||||
label: sbiPasskeys.label,
|
||||
keyringAccount: sbiPasskeys.keyringAccount,
|
||||
createdAt: sbiPasskeys.createdAt,
|
||||
updatedAt: sbiPasskeys.updatedAt,
|
||||
})
|
||||
.from(sbiPasskeys)
|
||||
.orderBy(sbiPasskeys.createdAt)
|
||||
const passkeys = await Promise.all(
|
||||
rows.map(async ({ keyringAccount, ...row }) => {
|
||||
const secret = await readSecret<StoredSbiPasskeySecret>(keyringAccount)
|
||||
const hasDeviceId = Boolean(effectiveSbiDeviceId(secret))
|
||||
const hasTradePassword = Boolean(effectiveSbiTradePassword(secret))
|
||||
return {
|
||||
...row,
|
||||
hasTradePassword,
|
||||
hasDeviceId,
|
||||
cashOrderReady: hasTradePassword && hasDeviceId,
|
||||
}
|
||||
}),
|
||||
)
|
||||
return textResult({ passkeys })
|
||||
},
|
||||
)
|
||||
|
||||
const callSbiMethod = async (method: RpcMethod, passkeyId: string, params: unknown) => {
|
||||
requireAuthenticated(auth)
|
||||
|
||||
if (auth.type === 'apiKey') {
|
||||
await assertApiKeyMethodAllowed(db, auth.apiKeyId, method)
|
||||
}
|
||||
|
||||
if (isTradingMethod(method)) {
|
||||
const tradingParams = params as { allowTrading?: boolean } | undefined
|
||||
if (!tradingParams?.allowTrading) throw new Error('trading methods require allowTrading')
|
||||
if (auth.type === 'apiKey') {
|
||||
await assertAndConsumeApiKeyTradeLimits({
|
||||
db,
|
||||
apiKeyId: auth.apiKeyId,
|
||||
params,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
if (isCashOrderMethod(method)) {
|
||||
const [passkey] = await db
|
||||
.select({ keyringAccount: sbiPasskeys.keyringAccount })
|
||||
.from(sbiPasskeys)
|
||||
.where(eq(sbiPasskeys.id, passkeyId))
|
||||
.limit(1)
|
||||
if (!passkey) throw new Error('SBI passkey not found')
|
||||
const secret = await readSecret<StoredSbiPasskeySecret>(passkey.keyringAccount)
|
||||
if (!effectiveSbiDeviceId(secret)) {
|
||||
throw new Error(
|
||||
'orders.cash methods require an SBI deviceId registered with F1131. This passkey has no saved deviceId, so MCP cannot complete SBI trade authentication for cash order estimates or orders.',
|
||||
)
|
||||
}
|
||||
if (!effectiveSbiTradePassword(secret)) {
|
||||
throw new Error(
|
||||
'orders.cash methods require a saved SBI tradePassword. This passkey has no saved tradePassword, so MCP cannot complete cash order estimates or orders.',
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
const client = await connectSbi(db, config, passkeyId)
|
||||
const result = await invokeSbiMethod(client, method, params)
|
||||
const submitMethod = submitMethodForEstimateMethod(method)
|
||||
if (!submitMethod) return textResult(result)
|
||||
|
||||
cleanupExpiredOrderSubmitTickets()
|
||||
const uuid = randomUUID()
|
||||
const expiresAt = new Date(Date.now() + ORDER_SUBMIT_TICKET_TTL_MS)
|
||||
const confirmationId = confirmationIdFromPreview(result)
|
||||
orderSubmitTickets.set(uuid, {
|
||||
passkeyId,
|
||||
estimateMethod: method,
|
||||
submitMethod,
|
||||
params,
|
||||
confirmationId,
|
||||
authKey: authKey(auth),
|
||||
expiresAt,
|
||||
})
|
||||
return textResult({
|
||||
preview: result,
|
||||
submit: {
|
||||
uuid,
|
||||
tool: 'csbie_sbi_submit_order',
|
||||
expiresAt: expiresAt.toISOString(),
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
server.registerTool(
|
||||
'csbie_sbi_submit_order',
|
||||
{
|
||||
title: 'Submit Estimated SBI Order',
|
||||
description:
|
||||
'Submit the same SBI order as a previous MCP estimate result by UUID. The UUID expires shortly and is bound to the same authenticated caller.',
|
||||
inputSchema: {
|
||||
uuid: z.string().uuid().describe('UUID returned by an order estimate tool'),
|
||||
},
|
||||
},
|
||||
async ({ uuid }) => {
|
||||
requireAuthenticated(auth)
|
||||
cleanupExpiredOrderSubmitTickets()
|
||||
|
||||
const ticket = orderSubmitTickets.get(uuid)
|
||||
if (!ticket) throw new Error('order submit uuid not found or expired')
|
||||
if (ticket.authKey !== authKey(auth)) {
|
||||
throw new Error('order submit uuid was created by a different authenticated caller')
|
||||
}
|
||||
|
||||
orderSubmitTickets.delete(uuid)
|
||||
return callSbiMethod(
|
||||
ticket.submitMethod,
|
||||
ticket.passkeyId,
|
||||
orderSubmitParams(ticket.params, ticket.confirmationId),
|
||||
)
|
||||
},
|
||||
)
|
||||
|
||||
for (const method of mcpExposedRpcMethods) {
|
||||
const paramsSchema = methodParamSchemas[method]
|
||||
server.registerTool(
|
||||
toolNameForMethod(method),
|
||||
{
|
||||
title: `Call ${method}`,
|
||||
description: `Connect with one saved SBI passkey and call ${method}. API key method permissions and trading limits are enforced.`,
|
||||
inputSchema: {
|
||||
passkeyId: z.string().describe('Saved SBI passkey id from csbie_sbi_passkeys'),
|
||||
...(paramsSchema
|
||||
? {
|
||||
params: paramsSchema.describe(`${method} params`),
|
||||
}
|
||||
: {}),
|
||||
},
|
||||
},
|
||||
async ({ passkeyId, params }) => callSbiMethod(method, passkeyId, params),
|
||||
)
|
||||
}
|
||||
|
||||
return server
|
||||
}
|
||||
|
||||
export const createMcpRoutes = () => {
|
||||
const app = new Hono<AppBindings>()
|
||||
|
||||
app.all('/', async (c) => {
|
||||
if (!c.get('authenticated')) {
|
||||
const resourceMetadata = new URL('/.well-known/oauth-protected-resource/api/mcp', c.req.url)
|
||||
c.header('WWW-Authenticate', `Bearer resource_metadata="${resourceMetadata.toString()}"`)
|
||||
return c.json({ error: 'unauthorized' }, 401)
|
||||
}
|
||||
|
||||
const transport = new StreamableHTTPTransport({
|
||||
sessionIdGenerator: undefined,
|
||||
enableJsonResponse: true,
|
||||
})
|
||||
const server = createMcpServer(c)
|
||||
|
||||
try {
|
||||
await server.connect(transport)
|
||||
return await transport.handleRequest(c)
|
||||
} finally {
|
||||
await server.close()
|
||||
await transport.close()
|
||||
}
|
||||
})
|
||||
|
||||
return app
|
||||
}
|
||||
@@ -0,0 +1,90 @@
|
||||
import { eq } from 'drizzle-orm'
|
||||
import { Hono } from 'hono'
|
||||
import type { AppBindings } from '../context'
|
||||
import { oauthClients } from '../db/schema'
|
||||
import type { ApiKeySettings } from '../security/api-keys'
|
||||
import { createOAuthAuthorizationCode } from '../security/oauth-provider'
|
||||
|
||||
const loopbackHosts = new Set(['localhost', '127.0.0.1', '[::1]'])
|
||||
|
||||
const redirectUriAllowed = (redirectUri: string, registeredUris: string[]) => {
|
||||
if (registeredUris.includes(redirectUri)) return true
|
||||
let requested: URL
|
||||
try {
|
||||
requested = new URL(redirectUri)
|
||||
} catch {
|
||||
return false
|
||||
}
|
||||
if (!loopbackHosts.has(requested.hostname)) return false
|
||||
return registeredUris.some((registeredUri) => {
|
||||
try {
|
||||
const registered = new URL(registeredUri)
|
||||
return (
|
||||
registered.protocol === requested.protocol &&
|
||||
registered.hostname === requested.hostname &&
|
||||
registered.pathname === requested.pathname &&
|
||||
registered.search === requested.search
|
||||
)
|
||||
} catch {
|
||||
return false
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
export const createOAuthRoutes = () => {
|
||||
const app = new Hono<AppBindings>()
|
||||
|
||||
app.get('/client/:id', async (c) => {
|
||||
if (!c.get('authenticated')) return c.json({ error: 'unauthorized' }, 401)
|
||||
const [client] = await c
|
||||
.get('db')
|
||||
.select()
|
||||
.from(oauthClients)
|
||||
.where(eq(oauthClients.id, c.req.param('id')))
|
||||
if (!client) return c.json({ error: 'client not found' }, 404)
|
||||
return c.json({ client: client.client })
|
||||
})
|
||||
|
||||
app.post('/approve', async (c) => {
|
||||
if (!c.get('authenticated')) return c.json({ error: 'unauthorized' }, 401)
|
||||
const body = await c.req.json<{
|
||||
clientId?: string
|
||||
redirectUri?: string
|
||||
codeChallenge?: string
|
||||
scope?: string
|
||||
state?: string
|
||||
resource?: string
|
||||
settings?: ApiKeySettings
|
||||
}>()
|
||||
if (!body.clientId || !body.redirectUri || !body.codeChallenge) {
|
||||
return c.json({ error: 'clientId, redirectUri and codeChallenge are required' }, 400)
|
||||
}
|
||||
|
||||
const [client] = await c
|
||||
.get('db')
|
||||
.select()
|
||||
.from(oauthClients)
|
||||
.where(eq(oauthClients.id, body.clientId))
|
||||
const redirectUris =
|
||||
(client?.client as { redirect_uris?: string[] } | undefined)?.redirect_uris ?? []
|
||||
if (!redirectUriAllowed(body.redirectUri, redirectUris)) {
|
||||
return c.json({ error: 'redirectUri is not registered for this client' }, 400)
|
||||
}
|
||||
|
||||
const code = await createOAuthAuthorizationCode(c.get('db'), {
|
||||
clientId: body.clientId,
|
||||
redirectUri: body.redirectUri,
|
||||
codeChallenge: body.codeChallenge,
|
||||
scopes: body.scope?.split(' ').filter(Boolean) ?? [],
|
||||
resource: body.resource,
|
||||
apiKeySettings: body.settings,
|
||||
})
|
||||
|
||||
const redirectUrl = new URL(body.redirectUri)
|
||||
redirectUrl.searchParams.set('code', code)
|
||||
if (body.state) redirectUrl.searchParams.set('state', body.state)
|
||||
return c.json({ redirectTo: redirectUrl.toString() })
|
||||
})
|
||||
|
||||
return app
|
||||
}
|
||||
@@ -0,0 +1,112 @@
|
||||
import type { SbiClientMethods } from '@repo/sbi-client'
|
||||
|
||||
export const RPC_METHODS = [
|
||||
'session.profile',
|
||||
'account.profile',
|
||||
'account.power.buyingPower',
|
||||
'account.power.collateralRatio',
|
||||
'account.positions.cash',
|
||||
'account.positions.cashDetail',
|
||||
'account.positions.cashForIssue',
|
||||
'account.positions.margin',
|
||||
'account.positions.marginDetail',
|
||||
'account.positions.marginForIssue',
|
||||
'account.positions.marginSummaryForIssue',
|
||||
'account.positions.marginDetailsForIssue',
|
||||
'account.positions.closeableMargin',
|
||||
'account.positions.deliverableMargin',
|
||||
'account.profitLoss.unrealized',
|
||||
'market.issue.search',
|
||||
'market.issue.suggest',
|
||||
'market.issue.allowedPrices',
|
||||
'market.issue.board',
|
||||
'market.issue.chart',
|
||||
'market.issue.openOrders',
|
||||
'market.issue.tradingInfo',
|
||||
'market.index.major',
|
||||
'market.overview',
|
||||
'market.ranking.market',
|
||||
'market.ranking.sector',
|
||||
'market.ranking.sbi',
|
||||
'news.list',
|
||||
'watchlist.list',
|
||||
'orders.inquiry.executionsToday',
|
||||
'orders.inquiry.open',
|
||||
'orders.cash.estimate',
|
||||
'orders.cash.place',
|
||||
'orders.cash.estimateCorrection',
|
||||
'orders.cash.estimateCorrectionConfirm',
|
||||
'orders.cash.placeCorrection',
|
||||
'orders.cash.estimateCancel',
|
||||
'orders.cash.placeCancel',
|
||||
'orders.margin.estimateOpen',
|
||||
'orders.margin.open',
|
||||
'orders.margin.estimateClose',
|
||||
'orders.margin.close',
|
||||
'orders.margin.estimateCloseSummary',
|
||||
'orders.margin.closeSummary',
|
||||
'orders.margin.estimateSummary',
|
||||
'orders.margin.placeSummary',
|
||||
'orders.margin.estimateActualDelivery',
|
||||
'orders.margin.actualDelivery',
|
||||
'orders.ifd.estimate',
|
||||
'orders.ifd.place',
|
||||
'orders.ifd.estimateCorrection',
|
||||
'orders.ifd.placeCorrection',
|
||||
'orders.ifd.estimateCancel',
|
||||
'orders.ifd.placeCancel',
|
||||
'orders.themeInvestment.list',
|
||||
'orders.themeInvestment.estimate',
|
||||
'orders.themeInvestment.place',
|
||||
] as const
|
||||
|
||||
export type RpcMethod = (typeof RPC_METHODS)[number]
|
||||
|
||||
const methodSet = new Set<string>(RPC_METHODS)
|
||||
const tradingMethods = new Set<string>([
|
||||
'orders.cash.place',
|
||||
'orders.cash.placeCorrection',
|
||||
'orders.cash.placeCancel',
|
||||
'orders.margin.open',
|
||||
'orders.margin.close',
|
||||
'orders.margin.closeSummary',
|
||||
'orders.margin.placeSummary',
|
||||
'orders.margin.actualDelivery',
|
||||
'orders.ifd.place',
|
||||
'orders.ifd.placeCorrection',
|
||||
'orders.ifd.placeCancel',
|
||||
'orders.themeInvestment.place',
|
||||
])
|
||||
|
||||
const cashOrderMethods = new Set<string>([
|
||||
'orders.cash.estimate',
|
||||
'orders.cash.place',
|
||||
'orders.cash.estimateCorrection',
|
||||
'orders.cash.estimateCorrectionConfirm',
|
||||
'orders.cash.placeCorrection',
|
||||
'orders.cash.estimateCancel',
|
||||
'orders.cash.placeCancel',
|
||||
])
|
||||
|
||||
export const isRpcMethod = (method: string): method is RpcMethod => methodSet.has(method)
|
||||
|
||||
export const isTradingMethod = (method: string) => tradingMethods.has(method)
|
||||
|
||||
export const isCashOrderMethod = (method: string) => cashOrderMethods.has(method)
|
||||
|
||||
export const invokeSbiMethod = async (
|
||||
client: SbiClientMethods,
|
||||
method: RpcMethod,
|
||||
params: unknown,
|
||||
) => {
|
||||
const target = method.split('.').reduce<unknown>((value, key) => {
|
||||
if (!value || typeof value !== 'object') return undefined
|
||||
return (value as Record<string, unknown>)[key]
|
||||
}, client)
|
||||
|
||||
if (typeof target !== 'function') throw new Error(`RPC method not callable: ${method}`)
|
||||
|
||||
if (Array.isArray(params)) return target(...params)
|
||||
if (params === undefined || params === null) return target()
|
||||
return target(params)
|
||||
}
|
||||
@@ -0,0 +1,33 @@
|
||||
import { eq } from 'drizzle-orm'
|
||||
import { loginWithPasskey } from '@repo/sbi-client'
|
||||
import type { SbiClientMethods, SbiClientOptions } from '@repo/sbi-client'
|
||||
import type { ServerConfig } from '../config'
|
||||
import type { Db } from '../db'
|
||||
import { sbiPasskeys } from '../db/schema'
|
||||
import type { StoredSbiPasskeySecret } from '../routes/admin'
|
||||
import { readSecret } from '../security/keyring'
|
||||
import { effectiveSbiDeviceId, effectiveSbiTradePassword } from '../security/sbi-credentials'
|
||||
|
||||
export const connectSbi = async (
|
||||
db: Db,
|
||||
config: ServerConfig,
|
||||
passkeyId: string,
|
||||
): Promise<SbiClientMethods> => {
|
||||
const [row] = await db.select().from(sbiPasskeys).where(eq(sbiPasskeys.id, passkeyId)).limit(1)
|
||||
if (!row) throw new Error('SBI passkey not found')
|
||||
|
||||
const secret = await readSecret<StoredSbiPasskeySecret>(row.keyringAccount)
|
||||
const clientOptions: SbiClientOptions = {
|
||||
tradePassword: effectiveSbiTradePassword(secret),
|
||||
deviceId: effectiveSbiDeviceId(secret),
|
||||
}
|
||||
return loginWithPasskey(
|
||||
{
|
||||
passkeyCredential: secret.credential,
|
||||
authBaseUrl: config.authBaseUrl,
|
||||
mtsBaseUrl: config.mtsBaseUrl,
|
||||
izanagiBaseUrl: config.izanagiBaseUrl,
|
||||
},
|
||||
clientOptions,
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,253 @@
|
||||
import type { SbiClientMethods } from '@repo/sbi-client'
|
||||
import { createBunWebSocket } from 'hono/bun'
|
||||
import type { WSContext } from 'hono/ws'
|
||||
import { randomUUID } from 'node:crypto'
|
||||
import type { ServerConfig } from '../config'
|
||||
import type { Db } from '../db'
|
||||
import {
|
||||
assertAndConsumeApiKeyTradeLimits,
|
||||
assertApiKeyMethodAllowed,
|
||||
} from '../security/trade-limits'
|
||||
import { invokeSbiMethod, isRpcMethod, isTradingMethod, RPC_METHODS } from './methods'
|
||||
import { connectSbi } from './sbi-session'
|
||||
|
||||
type JsonRpcRequest = {
|
||||
jsonrpc?: '2.0'
|
||||
id?: string | number | null
|
||||
method?: string
|
||||
params?: unknown
|
||||
}
|
||||
|
||||
type RpcSocketState = {
|
||||
client?: SbiClientMethods
|
||||
sbiPasskeyId?: string
|
||||
apiKeyId?: string
|
||||
boardPollingSubscriptions: Map<string, AbortController>
|
||||
}
|
||||
|
||||
type BoardPollingParams = {
|
||||
issueCode: string
|
||||
market?: string
|
||||
intervalSeconds?: number
|
||||
}
|
||||
|
||||
const BOARD_POLLING_METHODS = [
|
||||
'market.issue.pollBoard.subscribe',
|
||||
'market.issue.pollBoard.unsubscribe',
|
||||
] as const
|
||||
|
||||
const send = (ws: WSContext, payload: unknown) => {
|
||||
ws.send(JSON.stringify(payload))
|
||||
}
|
||||
|
||||
const result = (id: JsonRpcRequest['id'], value: unknown) => ({
|
||||
jsonrpc: '2.0',
|
||||
id: id ?? null,
|
||||
result: value,
|
||||
})
|
||||
|
||||
const error = (id: JsonRpcRequest['id'], code: number, message: string) => ({
|
||||
jsonrpc: '2.0',
|
||||
id: id ?? null,
|
||||
error: { code, message },
|
||||
})
|
||||
|
||||
const notification = (method: string, params: unknown) => ({
|
||||
jsonrpc: '2.0',
|
||||
method,
|
||||
params,
|
||||
})
|
||||
|
||||
const parseBoardPollingParams = (params: unknown): BoardPollingParams => {
|
||||
if (!params || typeof params !== 'object' || Array.isArray(params)) {
|
||||
throw new Error('pollBoard params are required')
|
||||
}
|
||||
|
||||
const value = params as Record<string, unknown>
|
||||
if (typeof value.issueCode !== 'string' || !value.issueCode) {
|
||||
throw new Error('issueCode is required')
|
||||
}
|
||||
if (value.market != null && typeof value.market !== 'string') {
|
||||
throw new Error('market must be a string')
|
||||
}
|
||||
if (
|
||||
value.intervalSeconds != null &&
|
||||
(typeof value.intervalSeconds !== 'number' ||
|
||||
!Number.isFinite(value.intervalSeconds) ||
|
||||
value.intervalSeconds <= 0)
|
||||
) {
|
||||
throw new Error('intervalSeconds must be a positive finite number')
|
||||
}
|
||||
|
||||
return {
|
||||
issueCode: value.issueCode,
|
||||
market: typeof value.market === 'string' ? value.market : undefined,
|
||||
intervalSeconds: typeof value.intervalSeconds === 'number' ? value.intervalSeconds : undefined,
|
||||
}
|
||||
}
|
||||
|
||||
const stopBoardPollingSubscription = (state: RpcSocketState, subscriptionId: string) => {
|
||||
const controller = state.boardPollingSubscriptions.get(subscriptionId)
|
||||
if (!controller) return false
|
||||
state.boardPollingSubscriptions.delete(subscriptionId)
|
||||
controller.abort(new Error('market issue board polling unsubscribed'))
|
||||
return true
|
||||
}
|
||||
|
||||
const stopBoardPollingSubscriptions = (state: RpcSocketState) => {
|
||||
for (const subscriptionId of state.boardPollingSubscriptions.keys()) {
|
||||
stopBoardPollingSubscription(state, subscriptionId)
|
||||
}
|
||||
}
|
||||
|
||||
const subscribeBoardPolling = async (
|
||||
db: Db,
|
||||
state: RpcSocketState,
|
||||
ws: WSContext,
|
||||
request: JsonRpcRequest,
|
||||
) => {
|
||||
if (!state.client) return error(request.id, 4001, 'SBI session is not connected')
|
||||
|
||||
if (state.apiKeyId) {
|
||||
await assertApiKeyMethodAllowed(db, state.apiKeyId, 'market.issue.board')
|
||||
}
|
||||
|
||||
const params = parseBoardPollingParams(request.params)
|
||||
const subscriptionId = randomUUID()
|
||||
const controller = new AbortController()
|
||||
state.boardPollingSubscriptions.set(subscriptionId, controller)
|
||||
|
||||
void (async () => {
|
||||
try {
|
||||
for await (const board of state.client!.market.issue.pollBoard({
|
||||
...params,
|
||||
signal: controller.signal,
|
||||
})) {
|
||||
if (!state.boardPollingSubscriptions.has(subscriptionId)) return
|
||||
send(ws, notification('market.issue.pollBoard.update', { subscriptionId, board }))
|
||||
}
|
||||
} catch (cause) {
|
||||
if (!controller.signal.aborted) {
|
||||
send(
|
||||
ws,
|
||||
notification('market.issue.pollBoard.error', {
|
||||
subscriptionId,
|
||||
message: cause instanceof Error ? cause.message : 'pollBoard failed',
|
||||
}),
|
||||
)
|
||||
}
|
||||
} finally {
|
||||
state.boardPollingSubscriptions.delete(subscriptionId)
|
||||
}
|
||||
})()
|
||||
|
||||
return result(request.id, { subscriptionId })
|
||||
}
|
||||
|
||||
const handleRpc = async (
|
||||
db: Db,
|
||||
config: ServerConfig,
|
||||
state: RpcSocketState,
|
||||
ws: WSContext,
|
||||
request: JsonRpcRequest,
|
||||
) => {
|
||||
if (request.method === 'rpc.methods') {
|
||||
return result(request.id, [...RPC_METHODS, ...BOARD_POLLING_METHODS])
|
||||
}
|
||||
|
||||
if (request.method === 'sbi.connect') {
|
||||
const passkeyId =
|
||||
request.params && typeof request.params === 'object'
|
||||
? (request.params as { passkeyId?: string }).passkeyId
|
||||
: undefined
|
||||
if (!passkeyId) throw new Error('passkeyId is required')
|
||||
stopBoardPollingSubscriptions(state)
|
||||
state.client = await connectSbi(db, config, passkeyId)
|
||||
state.sbiPasskeyId = passkeyId
|
||||
return result(request.id, { connected: true, passkeyId })
|
||||
}
|
||||
|
||||
if (request.method === 'market.issue.pollBoard.subscribe') {
|
||||
return subscribeBoardPolling(db, state, ws, request)
|
||||
}
|
||||
|
||||
if (request.method === 'market.issue.pollBoard.unsubscribe') {
|
||||
const subscriptionId =
|
||||
request.params && typeof request.params === 'object'
|
||||
? (request.params as { subscriptionId?: string }).subscriptionId
|
||||
: undefined
|
||||
if (!subscriptionId) throw new Error('subscriptionId is required')
|
||||
return result(request.id, {
|
||||
unsubscribed: stopBoardPollingSubscription(state, subscriptionId),
|
||||
})
|
||||
}
|
||||
|
||||
if (!request.method || !isRpcMethod(request.method)) {
|
||||
return error(request.id, -32601, 'method not found')
|
||||
}
|
||||
|
||||
if (!state.client) return error(request.id, 4001, 'SBI session is not connected')
|
||||
|
||||
if (state.apiKeyId) {
|
||||
await assertApiKeyMethodAllowed(db, state.apiKeyId, request.method)
|
||||
}
|
||||
|
||||
if (isTradingMethod(request.method)) {
|
||||
const params = request.params as { allowTrading?: boolean } | undefined
|
||||
if (!params?.allowTrading)
|
||||
return error(request.id, 4002, 'trading methods require allowTrading')
|
||||
if (state.apiKeyId) {
|
||||
await assertAndConsumeApiKeyTradeLimits({
|
||||
db,
|
||||
apiKeyId: state.apiKeyId,
|
||||
params: request.params,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
return result(request.id, await invokeSbiMethod(state.client, request.method, request.params))
|
||||
}
|
||||
|
||||
export const createRpcWebSocket = (db: Db, config: ServerConfig) => {
|
||||
const { upgradeWebSocket, websocket } = createBunWebSocket()
|
||||
|
||||
return {
|
||||
websocket,
|
||||
upgradeWebSocket: upgradeWebSocket(async (c) => {
|
||||
if (!c.get('authenticated')) {
|
||||
return {
|
||||
onOpen(_event, ws) {
|
||||
ws.close(1008, 'unauthorized')
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
const auth = c.get('auth')
|
||||
const state: RpcSocketState = {
|
||||
apiKeyId: auth.type === 'apiKey' ? auth.apiKeyId : undefined,
|
||||
boardPollingSubscriptions: new Map(),
|
||||
}
|
||||
|
||||
return {
|
||||
onOpen(_event, ws) {
|
||||
send(ws, result(null, { connected: true, methods: ['rpc.methods', 'sbi.connect'] }))
|
||||
},
|
||||
async onMessage(event, ws) {
|
||||
let request: JsonRpcRequest | undefined
|
||||
try {
|
||||
request = JSON.parse(String(event.data)) as JsonRpcRequest
|
||||
send(ws, await handleRpc(db, config, state, ws, request))
|
||||
} catch (cause) {
|
||||
send(
|
||||
ws,
|
||||
error(request?.id, -32603, cause instanceof Error ? cause.message : 'internal error'),
|
||||
)
|
||||
}
|
||||
},
|
||||
onClose() {
|
||||
stopBoardPollingSubscriptions(state)
|
||||
},
|
||||
}
|
||||
}),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
import { and, eq, isNull } from 'drizzle-orm'
|
||||
import type { Db } from '../db'
|
||||
import { apiKeys } from '../db/schema'
|
||||
import { randomId, randomToken, sha256 } from './crypto'
|
||||
|
||||
export type ApiKeySettings = {
|
||||
maxTradesPerHour?: number | null
|
||||
maxTradesPer6Hours?: number | null
|
||||
maxTradesPerDay?: number | null
|
||||
maxOrderPriceJpy?: number | null
|
||||
maxOrderAmountJpy?: number | null
|
||||
allowedMethods?: string[] | null
|
||||
}
|
||||
|
||||
const normalizeLimit = (value: unknown) => {
|
||||
if (value === undefined || value === null || value === '') return null
|
||||
const number = Number(value)
|
||||
if (!Number.isFinite(number) || number < 0) throw new Error('limit must be a positive number')
|
||||
return Math.floor(number)
|
||||
}
|
||||
|
||||
export const normalizeApiKeySettings = (settings: ApiKeySettings = {}) => ({
|
||||
maxTradesPerHour: normalizeLimit(settings.maxTradesPerHour),
|
||||
maxTradesPer6Hours: normalizeLimit(settings.maxTradesPer6Hours),
|
||||
maxTradesPerDay: normalizeLimit(settings.maxTradesPerDay),
|
||||
maxOrderPriceJpy: normalizeLimit(settings.maxOrderPriceJpy),
|
||||
maxOrderAmountJpy: normalizeLimit(settings.maxOrderAmountJpy),
|
||||
allowedMethods:
|
||||
settings.allowedMethods === undefined
|
||||
? null
|
||||
: settings.allowedMethods?.length
|
||||
? [...new Set(settings.allowedMethods)].sort()
|
||||
: null,
|
||||
})
|
||||
|
||||
export const listApiKeys = async (db: Db) =>
|
||||
db
|
||||
.select({
|
||||
id: apiKeys.id,
|
||||
label: apiKeys.label,
|
||||
maxTradesPerHour: apiKeys.maxTradesPerHour,
|
||||
maxTradesPer6Hours: apiKeys.maxTradesPer6Hours,
|
||||
maxTradesPerDay: apiKeys.maxTradesPerDay,
|
||||
maxOrderPriceJpy: apiKeys.maxOrderPriceJpy,
|
||||
maxOrderAmountJpy: apiKeys.maxOrderAmountJpy,
|
||||
allowedMethods: apiKeys.allowedMethods,
|
||||
createdAt: apiKeys.createdAt,
|
||||
lastUsedAt: apiKeys.lastUsedAt,
|
||||
revokedAt: apiKeys.revokedAt,
|
||||
})
|
||||
.from(apiKeys)
|
||||
.orderBy(apiKeys.createdAt)
|
||||
|
||||
export const createApiKey = async (db: Db, label: string, settings: ApiKeySettings = {}) => {
|
||||
const token = `csbie_${randomToken()}`
|
||||
const now = new Date()
|
||||
const row = {
|
||||
id: randomId('key'),
|
||||
label,
|
||||
tokenHash: sha256(token),
|
||||
...normalizeApiKeySettings(settings),
|
||||
createdAt: now,
|
||||
}
|
||||
await db.insert(apiKeys).values(row)
|
||||
return { ...row, token, tokenHash: undefined }
|
||||
}
|
||||
|
||||
export const updateApiKeySettings = async (db: Db, id: string, settings: ApiKeySettings) => {
|
||||
await db.update(apiKeys).set(normalizeApiKeySettings(settings)).where(eq(apiKeys.id, id))
|
||||
}
|
||||
|
||||
export const revokeApiKey = async (db: Db, id: string) => {
|
||||
await db.update(apiKeys).set({ revokedAt: new Date() }).where(eq(apiKeys.id, id))
|
||||
}
|
||||
|
||||
export const verifyApiKey = async (db: Db, token: string) => {
|
||||
const tokenHash = sha256(token)
|
||||
const [row] = await db
|
||||
.select()
|
||||
.from(apiKeys)
|
||||
.where(and(eq(apiKeys.tokenHash, tokenHash), isNull(apiKeys.revokedAt)))
|
||||
.limit(1)
|
||||
if (!row) return undefined
|
||||
await db.update(apiKeys).set({ lastUsedAt: new Date() }).where(eq(apiKeys.id, row.id))
|
||||
return row
|
||||
}
|
||||
@@ -0,0 +1,13 @@
|
||||
import { createHash, randomBytes, timingSafeEqual } from 'node:crypto'
|
||||
|
||||
export const randomId = (prefix: string) => `${prefix}_${randomBytes(18).toString('base64url')}`
|
||||
|
||||
export const randomToken = () => randomBytes(32).toString('base64url')
|
||||
|
||||
export const sha256 = (value: string) => createHash('sha256').update(value).digest('hex')
|
||||
|
||||
export const safeEqual = (left: string, right: string) => {
|
||||
const leftBuffer = Buffer.from(left)
|
||||
const rightBuffer = Buffer.from(right)
|
||||
return leftBuffer.length === rightBuffer.length && timingSafeEqual(leftBuffer, rightBuffer)
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
import { eq } from 'drizzle-orm'
|
||||
import type { ServerConfig } from '../config'
|
||||
import type { Db } from '../db'
|
||||
import { sessions } from '../db/schema'
|
||||
import { verifyApiKey } from './api-keys'
|
||||
import type { AuthContext } from '../context'
|
||||
|
||||
const readCookie = (header: string | null, name: string) => {
|
||||
if (!header) return undefined
|
||||
for (const part of header.split(';')) {
|
||||
const [key, ...value] = part.trim().split('=')
|
||||
if (key === name) return decodeURIComponent(value.join('='))
|
||||
}
|
||||
return undefined
|
||||
}
|
||||
|
||||
const readQueryApiKey = (request: Request) => {
|
||||
const key = new URL(request.url).searchParams.get('key')?.trim()
|
||||
return key || undefined
|
||||
}
|
||||
|
||||
const apiKeyAuth = async (db: Db, token: string) => {
|
||||
const apiKey = await verifyApiKey(db, token)
|
||||
return apiKey
|
||||
? ({ type: 'apiKey', authenticated: true, apiKeyId: apiKey.id } satisfies AuthContext)
|
||||
: ({ type: 'none', authenticated: false } satisfies AuthContext)
|
||||
}
|
||||
|
||||
export const authenticateRequest = async (db: Db, config: ServerConfig, request: Request) => {
|
||||
const authorization = request.headers.get('authorization')
|
||||
if (authorization?.startsWith('Bearer ')) {
|
||||
return apiKeyAuth(db, authorization.slice('Bearer '.length))
|
||||
}
|
||||
|
||||
const queryApiKey = readQueryApiKey(request)
|
||||
if (queryApiKey) {
|
||||
return apiKeyAuth(db, queryApiKey)
|
||||
}
|
||||
|
||||
const sessionId = readCookie(request.headers.get('cookie'), config.sessionCookieName)
|
||||
if (!sessionId) return { type: 'none', authenticated: false } satisfies AuthContext
|
||||
const [session] = await db.select().from(sessions).where(eq(sessions.id, sessionId)).limit(1)
|
||||
if (!session || session.expiresAt <= new Date()) {
|
||||
return { type: 'none', authenticated: false } satisfies AuthContext
|
||||
}
|
||||
return { type: 'session', authenticated: true, sessionId } satisfies AuthContext
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
import { getPassword, setPassword, deletePassword } from '@napi-rs/keyring/keytar'
|
||||
|
||||
const SERVICE = 'csbie'
|
||||
|
||||
export const saveSecret = async (account: string, secret: unknown) => {
|
||||
await setPassword(SERVICE, account, JSON.stringify(secret))
|
||||
}
|
||||
|
||||
export const readSecret = async <T>(account: string): Promise<T> => {
|
||||
const secret = await getPassword(SERVICE, account)
|
||||
if (!secret) throw new Error(`secret not found: ${account}`)
|
||||
return JSON.parse(secret) as T
|
||||
}
|
||||
|
||||
export const deleteSecret = async (account: string) => {
|
||||
await deletePassword(SERVICE, account)
|
||||
}
|
||||
@@ -0,0 +1,191 @@
|
||||
import { and, eq, isNull } from 'drizzle-orm'
|
||||
import type { Context } from 'hono'
|
||||
import type { OAuthServerProvider } from '@modelcontextprotocol/sdk/server/auth/provider.js'
|
||||
import type {
|
||||
OAuthClientInformationFull,
|
||||
OAuthTokens,
|
||||
} from '@modelcontextprotocol/sdk/shared/auth.js'
|
||||
import type { AuthInfo } from '@modelcontextprotocol/sdk/server/auth/types.js'
|
||||
import type { ServerConfig } from '../config'
|
||||
import type { AppBindings } from '../context'
|
||||
import type { Db } from '../db'
|
||||
import { apiKeys, oauthAuthorizationCodes, oauthClients, oauthRefreshTokens } from '../db/schema'
|
||||
import { createApiKey, type ApiKeySettings, verifyApiKey } from './api-keys'
|
||||
import { randomToken, sha256 } from './crypto'
|
||||
|
||||
const CODE_TTL_MS = 10 * 60 * 1000
|
||||
const ACCESS_TOKEN_TTL_SECONDS = 60 * 60
|
||||
const REFRESH_TOKEN_TTL_MS = 30 * 24 * 60 * 60 * 1000
|
||||
|
||||
export const createOAuthAuthorizationCode = async (
|
||||
db: Db,
|
||||
options: {
|
||||
clientId: string
|
||||
redirectUri: string
|
||||
codeChallenge: string
|
||||
scopes: string[]
|
||||
resource?: string
|
||||
apiKeySettings?: ApiKeySettings
|
||||
},
|
||||
) => {
|
||||
const now = new Date()
|
||||
const code = `mcp_code_${randomToken()}`
|
||||
await db.insert(oauthAuthorizationCodes).values({
|
||||
code,
|
||||
clientId: options.clientId,
|
||||
redirectUri: options.redirectUri,
|
||||
codeChallenge: options.codeChallenge,
|
||||
scopes: options.scopes,
|
||||
resource: options.resource,
|
||||
apiKeySettings: (options.apiKeySettings ?? null) as Record<string, unknown> | null,
|
||||
createdAt: now,
|
||||
expiresAt: new Date(now.getTime() + CODE_TTL_MS),
|
||||
})
|
||||
return code
|
||||
}
|
||||
|
||||
export const createOAuthServerProvider = (db: Db, config: ServerConfig): OAuthServerProvider => ({
|
||||
get clientsStore() {
|
||||
return {
|
||||
getClient: async (clientId: string) => {
|
||||
const [row] = await db.select().from(oauthClients).where(eq(oauthClients.id, clientId))
|
||||
return row?.client as OAuthClientInformationFull | undefined
|
||||
},
|
||||
registerClient: async (
|
||||
client: Omit<OAuthClientInformationFull, 'client_id' | 'client_id_issued_at'>,
|
||||
) => {
|
||||
const clientInfo = client as OAuthClientInformationFull
|
||||
await db
|
||||
.insert(oauthClients)
|
||||
.values({
|
||||
id: clientInfo.client_id,
|
||||
client: clientInfo as unknown as Record<string, unknown>,
|
||||
createdAt: new Date(),
|
||||
})
|
||||
.onConflictDoUpdate({
|
||||
target: oauthClients.id,
|
||||
set: { client: clientInfo as unknown as Record<string, unknown> },
|
||||
})
|
||||
return clientInfo
|
||||
},
|
||||
}
|
||||
},
|
||||
|
||||
authorize: async (client, params, c: Context<AppBindings>) => {
|
||||
const approvalUrl = new URL('/oauth/authorize', config.origin)
|
||||
approvalUrl.searchParams.set('client_id', client.client_id)
|
||||
approvalUrl.searchParams.set('redirect_uri', params.redirectUri)
|
||||
approvalUrl.searchParams.set('code_challenge', params.codeChallenge)
|
||||
if (params.state) approvalUrl.searchParams.set('state', params.state)
|
||||
if (params.scopes?.length) approvalUrl.searchParams.set('scope', params.scopes.join(' '))
|
||||
if (params.resource) approvalUrl.searchParams.set('resource', params.resource.href)
|
||||
|
||||
const auth = c.get('auth')
|
||||
if (!auth?.authenticated || auth.type !== 'session') {
|
||||
approvalUrl.searchParams.set('login_required', '1')
|
||||
c.res = c.redirect(approvalUrl.toString(), 302)
|
||||
return
|
||||
}
|
||||
|
||||
c.res = c.redirect(approvalUrl.toString(), 302)
|
||||
},
|
||||
|
||||
challengeForAuthorizationCode: async (_client, authorizationCode) => {
|
||||
const [row] = await db
|
||||
.select()
|
||||
.from(oauthAuthorizationCodes)
|
||||
.where(eq(oauthAuthorizationCodes.code, authorizationCode))
|
||||
if (!row || row.expiresAt <= new Date()) throw new Error('authorization code expired')
|
||||
return row.codeChallenge
|
||||
},
|
||||
|
||||
exchangeAuthorizationCode: async (
|
||||
client,
|
||||
authorizationCode,
|
||||
_codeVerifier,
|
||||
redirectUri,
|
||||
resource,
|
||||
) => {
|
||||
const [code] = await db
|
||||
.select()
|
||||
.from(oauthAuthorizationCodes)
|
||||
.where(eq(oauthAuthorizationCodes.code, authorizationCode))
|
||||
if (!code || code.expiresAt <= new Date()) throw new Error('authorization code expired')
|
||||
if (code.clientId !== client.client_id) throw new Error('authorization code client mismatch')
|
||||
if (redirectUri && code.redirectUri !== redirectUri) throw new Error('redirect_uri mismatch')
|
||||
await db
|
||||
.delete(oauthAuthorizationCodes)
|
||||
.where(eq(oauthAuthorizationCodes.code, authorizationCode))
|
||||
|
||||
const key = await createApiKey(
|
||||
db,
|
||||
`OAuth: ${client.client_name ?? client.client_id}`,
|
||||
(code.apiKeySettings ?? {}) as ApiKeySettings,
|
||||
)
|
||||
const refreshToken = `csbie_refresh_${randomToken()}`
|
||||
const now = new Date()
|
||||
const scopes = code.scopes
|
||||
await db.insert(oauthRefreshTokens).values({
|
||||
tokenHash: sha256(refreshToken),
|
||||
clientId: client.client_id,
|
||||
apiKeyId: key.id,
|
||||
scopes,
|
||||
resource: resource?.href ?? code.resource,
|
||||
createdAt: now,
|
||||
expiresAt: new Date(now.getTime() + REFRESH_TOKEN_TTL_MS),
|
||||
})
|
||||
|
||||
return {
|
||||
access_token: key.token,
|
||||
token_type: 'Bearer',
|
||||
expires_in: ACCESS_TOKEN_TTL_SECONDS,
|
||||
refresh_token: refreshToken,
|
||||
scope: scopes.join(' '),
|
||||
} satisfies OAuthTokens
|
||||
},
|
||||
|
||||
exchangeRefreshToken: async (client, refreshToken, scopes, _resource) => {
|
||||
const [token] = await db
|
||||
.select()
|
||||
.from(oauthRefreshTokens)
|
||||
.where(
|
||||
and(
|
||||
eq(oauthRefreshTokens.tokenHash, sha256(refreshToken)),
|
||||
eq(oauthRefreshTokens.clientId, client.client_id),
|
||||
isNull(oauthRefreshTokens.revokedAt),
|
||||
),
|
||||
)
|
||||
if (!token || token.expiresAt <= new Date()) throw new Error('refresh token expired')
|
||||
|
||||
const newAccess = await createApiKey(db, `OAuth: ${client.client_name ?? client.client_id}`)
|
||||
return {
|
||||
access_token: newAccess.token,
|
||||
token_type: 'Bearer',
|
||||
expires_in: ACCESS_TOKEN_TTL_SECONDS,
|
||||
scope: (scopes ?? token.scopes).join(' '),
|
||||
refresh_token: refreshToken,
|
||||
} satisfies OAuthTokens
|
||||
},
|
||||
|
||||
verifyAccessToken: async (token) => {
|
||||
const apiKey = await verifyApiKey(db, token)
|
||||
if (!apiKey) throw new Error('invalid access token')
|
||||
return {
|
||||
token,
|
||||
clientId: apiKey.id,
|
||||
scopes: ['mcp'],
|
||||
expiresAt: Math.floor(Date.now() / 1000) + ACCESS_TOKEN_TTL_SECONDS,
|
||||
} satisfies AuthInfo
|
||||
},
|
||||
|
||||
revokeToken: async (_client, request) => {
|
||||
await db
|
||||
.update(apiKeys)
|
||||
.set({ revokedAt: new Date() })
|
||||
.where(eq(apiKeys.tokenHash, sha256(request.token)))
|
||||
await db
|
||||
.update(oauthRefreshTokens)
|
||||
.set({ revokedAt: new Date() })
|
||||
.where(eq(oauthRefreshTokens.tokenHash, sha256(request.token)))
|
||||
},
|
||||
})
|
||||
@@ -0,0 +1,11 @@
|
||||
import type { StoredSbiPasskeySecret } from '../routes/admin'
|
||||
|
||||
const nonEmpty = (value: string | undefined) => value?.trim() || undefined
|
||||
|
||||
export const effectiveSbiDeviceId = (secret: StoredSbiPasskeySecret) =>
|
||||
nonEmpty(secret.deviceId) ?? nonEmpty(process.env.SBI_DEVICE_ID)
|
||||
|
||||
export const effectiveSbiTradePassword = (secret: StoredSbiPasskeySecret) =>
|
||||
nonEmpty(secret.tradePassword) ?? nonEmpty(process.env.SBI_TRADE_PASSWORD)
|
||||
|
||||
export const hasNonEmptySecretValue = (value: string | undefined) => Boolean(nonEmpty(value))
|
||||
@@ -0,0 +1,44 @@
|
||||
import { eq, lt } from 'drizzle-orm'
|
||||
import type { Context } from 'hono'
|
||||
import { deleteCookie, getCookie, setCookie } from 'hono/cookie'
|
||||
import type { ServerConfig } from '../config'
|
||||
import type { Db } from '../db'
|
||||
import { sessions } from '../db/schema'
|
||||
import { randomId } from './crypto'
|
||||
|
||||
const SESSION_DAYS = 30
|
||||
|
||||
export const createSession = async (db: Db) => {
|
||||
const now = new Date()
|
||||
const expiresAt = new Date(now.getTime() + SESSION_DAYS * 24 * 60 * 60 * 1000)
|
||||
const id = randomId('ses')
|
||||
await db.insert(sessions).values({ id, createdAt: now, expiresAt })
|
||||
return { id, expiresAt }
|
||||
}
|
||||
|
||||
export const setSessionCookie = (c: Context, config: ServerConfig, id: string, expires: Date) => {
|
||||
setCookie(c, config.sessionCookieName, id, {
|
||||
httpOnly: true,
|
||||
sameSite: 'Lax',
|
||||
secure: config.origin.startsWith('https://'),
|
||||
path: '/',
|
||||
expires,
|
||||
})
|
||||
}
|
||||
|
||||
export const clearSessionCookie = async (c: Context, db: Db, config: ServerConfig) => {
|
||||
const id = getCookie(c, config.sessionCookieName)
|
||||
if (id) await db.delete(sessions).where(eq(sessions.id, id))
|
||||
deleteCookie(c, config.sessionCookieName, { path: '/' })
|
||||
}
|
||||
|
||||
export const verifySessionCookie = async (c: Context, db: Db, config: ServerConfig) => {
|
||||
const id = getCookie(c, config.sessionCookieName)
|
||||
if (!id) return false
|
||||
const [session] = await db.select().from(sessions).where(eq(sessions.id, id)).limit(1)
|
||||
return Boolean(session && session.expiresAt > new Date())
|
||||
}
|
||||
|
||||
export const pruneSessions = async (db: Db) => {
|
||||
await db.delete(sessions).where(lt(sessions.expiresAt, new Date()))
|
||||
}
|
||||
@@ -0,0 +1,11 @@
|
||||
import { sha256, safeEqual } from './crypto'
|
||||
|
||||
export const verifySetupPassword = (password: string) => {
|
||||
const hash = process.env.CSBIE_SETUP_PASSWORD_HASH
|
||||
if (hash) return safeEqual(sha256(password), hash)
|
||||
|
||||
const plaintext = process.env.CSBIE_SETUP_PASSWORD
|
||||
if (plaintext) return safeEqual(password, plaintext)
|
||||
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,117 @@
|
||||
import { and, eq } from 'drizzle-orm'
|
||||
import type { Db } from '../db'
|
||||
import { apiKeyTradeUsage, apiKeys } from '../db/schema'
|
||||
import type { RpcMethod } from '../rpc/methods'
|
||||
|
||||
type WindowName = '1h' | '3h' | '1d'
|
||||
|
||||
type TradeLimitInput = {
|
||||
apiKeyId: string
|
||||
params: unknown
|
||||
}
|
||||
|
||||
const windowBucket = (date: Date, window: WindowName) => {
|
||||
const year = date.getUTCFullYear()
|
||||
const month = String(date.getUTCMonth() + 1).padStart(2, '0')
|
||||
const day = String(date.getUTCDate()).padStart(2, '0')
|
||||
if (window === '1d') return `${year}-${month}-${day}`
|
||||
|
||||
const hour = date.getUTCHours()
|
||||
const bucketHour = window === '3h' ? Math.floor(hour / 3) * 3 : hour
|
||||
return `${year}-${month}-${day}T${String(bucketHour).padStart(2, '0')}`
|
||||
}
|
||||
|
||||
const numberFromParams = (params: unknown, key: string) => {
|
||||
if (!params || typeof params !== 'object' || Array.isArray(params)) return undefined
|
||||
const value = (params as Record<string, unknown>)[key]
|
||||
if (value === undefined || value === null || value === '') return undefined
|
||||
const number = Number(value)
|
||||
return Number.isFinite(number) ? number : undefined
|
||||
}
|
||||
|
||||
const assertPriceLimits = (
|
||||
params: unknown,
|
||||
settings: {
|
||||
maxOrderPriceJpy: number | null
|
||||
maxOrderAmountJpy: number | null
|
||||
},
|
||||
) => {
|
||||
const price = numberFromParams(params, 'price')
|
||||
const quantity = numberFromParams(params, 'quantity')
|
||||
|
||||
if (settings.maxOrderPriceJpy != null && price != null && price > settings.maxOrderPriceJpy) {
|
||||
throw new Error(`order price exceeds API key limit (${settings.maxOrderPriceJpy} JPY)`)
|
||||
}
|
||||
|
||||
if (settings.maxOrderAmountJpy == null) return
|
||||
if (price == null || quantity == null) {
|
||||
throw new Error('order amount limit requires price and quantity in trading params')
|
||||
}
|
||||
|
||||
const amount = price * quantity
|
||||
if (amount > settings.maxOrderAmountJpy) {
|
||||
throw new Error(`order amount exceeds API key limit (${settings.maxOrderAmountJpy} JPY)`)
|
||||
}
|
||||
}
|
||||
|
||||
const checkAndConsumeWindow = async (
|
||||
db: Db,
|
||||
apiKeyId: string,
|
||||
window: WindowName,
|
||||
limit: number | null,
|
||||
now: Date,
|
||||
) => {
|
||||
if (limit == null) return
|
||||
const hourBucket = windowBucket(now, window)
|
||||
const where = and(
|
||||
eq(apiKeyTradeUsage.apiKeyId, apiKeyId),
|
||||
eq(apiKeyTradeUsage.window, window),
|
||||
eq(apiKeyTradeUsage.hourBucket, hourBucket),
|
||||
)
|
||||
const [usage] = await db.select().from(apiKeyTradeUsage).where(where).limit(1)
|
||||
const currentCount = usage?.tradeCount ?? 0
|
||||
if (currentCount >= limit) throw new Error(`${window} trade limit exceeded for API key`)
|
||||
|
||||
if (usage) {
|
||||
await db
|
||||
.update(apiKeyTradeUsage)
|
||||
.set({ tradeCount: currentCount + 1, updatedAt: now })
|
||||
.where(where)
|
||||
return
|
||||
}
|
||||
|
||||
await db.insert(apiKeyTradeUsage).values({
|
||||
apiKeyId,
|
||||
window,
|
||||
hourBucket,
|
||||
tradeCount: 1,
|
||||
updatedAt: now,
|
||||
})
|
||||
}
|
||||
|
||||
export const assertAndConsumeApiKeyTradeLimits = async ({
|
||||
db,
|
||||
apiKeyId,
|
||||
params,
|
||||
}: TradeLimitInput & { db: Db }) => {
|
||||
const [key] = await db.select().from(apiKeys).where(eq(apiKeys.id, apiKeyId)).limit(1)
|
||||
if (!key || key.revokedAt) throw new Error('API key is not active')
|
||||
|
||||
assertPriceLimits(params, {
|
||||
maxOrderPriceJpy: key.maxOrderPriceJpy,
|
||||
maxOrderAmountJpy: key.maxOrderAmountJpy,
|
||||
})
|
||||
|
||||
const now = new Date()
|
||||
await checkAndConsumeWindow(db, apiKeyId, '1h', key.maxTradesPerHour, now)
|
||||
await checkAndConsumeWindow(db, apiKeyId, '3h', key.maxTradesPer6Hours, now)
|
||||
await checkAndConsumeWindow(db, apiKeyId, '1d', key.maxTradesPerDay, now)
|
||||
}
|
||||
|
||||
export const assertApiKeyMethodAllowed = async (db: Db, apiKeyId: string, method: RpcMethod) => {
|
||||
const [key] = await db.select().from(apiKeys).where(eq(apiKeys.id, apiKeyId)).limit(1)
|
||||
if (!key || key.revokedAt) throw new Error('API key is not active')
|
||||
if (key.allowedMethods === null || key.allowedMethods === undefined) return
|
||||
if (key.allowedMethods.includes(method)) return
|
||||
throw new Error(`API key is not allowed to call ${method}`)
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
{
|
||||
"extends": "../../tsconfig.json",
|
||||
"compilerOptions": {
|
||||
"types": ["bun"],
|
||||
"rootDir": "src",
|
||||
"outDir": "dist",
|
||||
"noEmit": true
|
||||
},
|
||||
"include": ["src/**/*.ts"]
|
||||
}
|
||||
Reference in New Issue
Block a user