refactor(server): enhance validation schemas for auth and audio processing

This commit is contained in:
RainbowBird
2026-06-03 13:48:19 +08:00
parent a795806e81
commit e4acc78e25
4 changed files with 47 additions and 30 deletions
@@ -7,6 +7,13 @@ import { createHmac } from 'node:crypto'
import { createAuthMiddleware } from 'better-auth/api'
import { createLocalJWKSet, jwtVerify } from 'jose'
import { pipe, regex, safeParse, string, transform } from 'valibot'
const JwtBearerTokenSchema = pipe(
string(),
transform(value => value.trim()),
regex(/^[\w-]+\.[\w-]+\.[\w-]+$/, 'Bearer token must be a compact JWT'),
)
/**
* Bridge plugin that lets better-auth's `sessionMiddleware` accept the
@@ -64,11 +71,6 @@ import { createLocalJWKSet, jwtVerify } from 'jose'
* externally-signed JWTs against a JWKS for its own session resolution.
*/
export function oidcJwtBearer(env: Env): BetterAuthPlugin {
// JWT shape: three base64url segments separated by dots. Catches the
// happy path without us decoding; downstream JWKS verify is the real
// gate. Anything that fails this regex falls through to bearer().
const JWT_SHAPE_RE = /^[\w-]+\.[\w-]+\.[\w-]+$/
// Bridge session lifetime. Long enough to span an OAuth round-trip
// (link-social → provider → callback) on slow networks; short enough
// that an unused row TTL-prunes quickly.
@@ -207,9 +209,14 @@ export function oidcJwtBearer(env: Env): BetterAuthPlugin {
if (lower !== 'bearer ')
return
const token = authHeader.slice(7).trim()
if (!token || !JWT_SHAPE_RE.test(token))
// JWT shape: three base64url segments separated by dots. Catches
// the happy path without us decoding; downstream JWKS verify is
// the real gate. Anything that fails this schema falls through to
// bearer().
const tokenResult = safeParse(JwtBearerTokenSchema, authHeader.slice(7))
if (!tokenResult.success)
return
const token = tokenResult.output
// Verify against our own JWKS, read directly from DB (no
// self-fetch). If it isn't ours (signature mismatch, wrong
@@ -1,5 +1,19 @@
import type { RawData } from 'ws'
import { Buffer } from 'node:buffer'
import { finite, looseObject, minValue, number, optional, pipe, safeParse } from 'valibot'
const UpstreamUsagePayloadSchema = looseObject({
usage: optional(looseObject({
text_words: optional(pipe(
number(),
finite(),
minValue(0),
)),
})),
})
/**
* Normalizes websocket text payload chunks.
*
@@ -11,7 +25,7 @@ import { Buffer } from 'node:buffer'
* - `"frame"`
* - `"ab"`
*/
export function bufferToString(data: Buffer | Buffer[] | ArrayBuffer): string {
export function bufferToString(data: RawData): string {
if (Array.isArray(data))
return Buffer.concat(data).toString('utf8')
if (data instanceof ArrayBuffer)
@@ -29,7 +43,7 @@ export function bufferToString(data: Buffer | Buffer[] | ArrayBuffer): string {
* After:
* - `ArrayBuffer`
*/
export function toBufferLike(data: Buffer | Buffer[] | ArrayBuffer): ArrayBuffer {
export function toBufferLike(data: RawData): ArrayBuffer {
if (Array.isArray(data)) {
const merged = Buffer.concat(data)
return merged.buffer.slice(merged.byteOffset, merged.byteOffset + merged.byteLength) as ArrayBuffer
@@ -51,13 +65,7 @@ export function toBufferLike(data: Buffer | Buffer[] | ArrayBuffer): ArrayBuffer
* - `null`
*/
export function readUsageChars(payload: Record<string, unknown> | undefined): number | null {
if (!payload || typeof payload !== 'object')
return null
const usage = (payload as { usage?: unknown }).usage
if (!usage || typeof usage !== 'object')
return null
const textWords = (usage as { text_words?: unknown }).text_words
if (typeof textWords === 'number' && Number.isFinite(textWords) && textWords >= 0)
return Math.floor(textWords)
return null
const result = safeParse(UpstreamUsagePayloadSchema, payload)
const textWords = result.success ? result.output.usage?.text_words : undefined
return typeof textWords === 'number' ? Math.floor(textWords) : null
}
@@ -1,4 +1,5 @@
import type { WSContext } from 'hono/ws'
import type { RawData } from 'ws'
import type { FluxService } from '../../services/domain/flux'
import type { AudioSpeechWsHandlersOptions } from './types'
@@ -253,7 +254,7 @@ export function createSessionState(userId: string, opts: AudioSpeechWsHandlersOp
finalize()
}
function handleUpstreamMessage(data: Buffer | Buffer[] | ArrayBuffer, isBinary: boolean) {
function handleUpstreamMessage(data: RawData, isBinary: boolean) {
if (!clientWs)
return
if (isBinary) {
+12 -11
View File
@@ -1,16 +1,19 @@
import type { Database } from '../../libs/db'
import { and, eq } from 'drizzle-orm'
import { email, nonEmpty, object, pipe, safeParse, string, transform } from 'valibot'
import { account, user } from '../../schemas/accounts'
import { createBadRequestError } from '../../utils/error'
// NOTICE:
// Loose RFC-5322-ish regex used to fail fast on obviously malformed input.
// Authoritative validation happens in better-auth on sign-in/sign-up;
// this is just a pre-flight gate for the email-first identifier step so we
// avoid hitting the DB with garbage.
const EMAIL_SHAPE_RE = /^[^\s@]+@[^\s@][^\s.@]*\.[^\s@]+$/
const CheckEmailIdentifierBodySchema = object({
email: pipe(
string(),
transform(value => value.trim().toLowerCase()),
nonEmpty('email is required'),
email('email must be a valid email address'),
),
})
export interface CheckEmailIdentifierDeps {
/** Database used to inspect user and credential-account rows. */
@@ -41,16 +44,14 @@ export async function checkEmailIdentifier(
deps: CheckEmailIdentifierDeps,
body: { email?: unknown } | null,
): Promise<CheckEmailIdentifierResult> {
const raw = typeof body?.email === 'string' ? body.email.trim() : ''
const email = raw.toLowerCase()
if (!email || !EMAIL_SHAPE_RE.test(email))
const parsed = safeParse(CheckEmailIdentifierBodySchema, body)
if (!parsed.success)
throw createBadRequestError('Invalid email', 'INVALID_EMAIL')
const [matched] = await deps.db
.select({ id: user.id })
.from(user)
.where(eq(user.email, email))
.where(eq(user.email, parsed.output.email))
.limit(1)
if (!matched)