import { and, eq, gt, sql } from 'drizzle-orm'
import type { Querier } from '@platform-modules/db'
import { generateOpaqueToken, hashToken } from '@platform-modules/util/tokens'
import { authUsers, userSessions } from './schema.js'
import { signSession, ACCESS_TOKEN_TTL_SECS } from './session.js'

export const SESSION_TTL_MS = 30 * 24 * 60 * 60 * 1000
export const GRACE_WINDOW_MS = 60_000

export type SessionUserRow = {
  sessionId: string
  userId: string
  sessionVersion: number
  roles: string[]
}

export type RotateRefreshResult =
  | {
      kind: 'rotated'
      user: SessionUserRow
      accessToken: string
      refreshToken: string
    }
  | {
      kind: 'grace'
      user: SessionUserRow
      accessToken: string
    }
  | {
      kind: 'race_lost'
      user: SessionUserRow
      accessToken: string
    }
  | {
      kind: 'reuse_detected'
      userId: string
      sessionId: string
      newSessionVersion: number
    }
  | { kind: 'invalid' }

function parseRoles(raw: string): string[] {
  try {
    const parsed = JSON.parse(raw) as unknown
    return Array.isArray(parsed) ? (parsed as string[]) : ['user']
  } catch {
    return ['user']
  }
}

async function mintAccessToken(
  user: SessionUserRow,
  jwtSecrets: string[],
): Promise<string> {
  return signSession(
    {
      sub: user.userId,
      sid: user.sessionId,
      sv: user.sessionVersion,
    },
    jwtSecrets[0]!,
    ACCESS_TOKEN_TTL_SECS,
  )
}

export async function rotateRefreshToken<S extends Record<string, unknown>>(
  db: Querier<S>,
  tables: { authUsers: typeof authUsers; userSessions: typeof userSessions },
  rawRefreshToken: string,
  jwtSecrets: string[],
): Promise<RotateRefreshResult> {
  const oldHash = await hashToken(rawRefreshToken)
  const now = new Date()

  const [row] = await db
    .select({
      sessionId: tables.userSessions.id,
      userId: tables.userSessions.userId,
      sessionVersion: tables.authUsers.sessionVersion,
      roles: tables.authUsers.roles,
      status: tables.userSessions.status,
    })
    .from(tables.userSessions)
    .innerJoin(tables.authUsers, eq(tables.userSessions.userId, tables.authUsers.id))
    .where(
      and(
        eq(tables.userSessions.refreshTokenHash, oldHash),
        eq(tables.userSessions.status, 'active'),
        gt(tables.userSessions.expiresAt, now),
      ),
    )
    .limit(1)

  if (!row) {
    const graceCutoff = new Date(now.getTime() - GRACE_WINDOW_MS)
    const [graceRow] = await db
      .select({
        sessionId: tables.userSessions.id,
        userId: tables.userSessions.userId,
        sessionVersion: tables.authUsers.sessionVersion,
        roles: tables.authUsers.roles,
      })
      .from(tables.userSessions)
      .innerJoin(tables.authUsers, eq(tables.userSessions.userId, tables.authUsers.id))
      .where(
        and(
          eq(tables.userSessions.previousRefreshTokenHash, oldHash),
          gt(tables.userSessions.lastRefreshedAt, graceCutoff),
          eq(tables.userSessions.status, 'active'),
          gt(tables.userSessions.expiresAt, now),
        ),
      )
      .limit(1)

    if (graceRow) {
      const user: SessionUserRow = {
        sessionId: graceRow.sessionId,
        userId: graceRow.userId,
        sessionVersion: graceRow.sessionVersion,
        roles: parseRoles(graceRow.roles),
      }
      return { kind: 'grace', user, accessToken: await mintAccessToken(user, jwtSecrets) }
    }

    // P2 — consumed RT replay past grace → family revoke (RFC 6819 §5.2.2.3).
    const [reuseRow] = await db
      .select({
        sessionId: tables.userSessions.id,
        userId: tables.userSessions.userId,
      })
      .from(tables.userSessions)
      .where(eq(tables.userSessions.previousRefreshTokenHash, oldHash))
      .limit(1)

    if (reuseRow) {
      await db
        .update(tables.userSessions)
        .set({ status: 'revoked' })
        .where(eq(tables.userSessions.id, reuseRow.sessionId))

      const [bumped] = await db
        .update(tables.authUsers)
        .set({ sessionVersion: sql`${tables.authUsers.sessionVersion} + 1` })
        .where(eq(tables.authUsers.id, reuseRow.userId))
        .returning({ sessionVersion: tables.authUsers.sessionVersion })

      return {
        kind: 'reuse_detected',
        userId: reuseRow.userId,
        sessionId: reuseRow.sessionId,
        newSessionVersion: bumped?.sessionVersion ?? 0,
      }
    }

    return { kind: 'invalid' }
  }

  const user: SessionUserRow = {
    sessionId: row.sessionId,
    userId: row.userId,
    sessionVersion: row.sessionVersion,
    roles: parseRoles(row.roles),
  }

  const newRefreshToken = generateOpaqueToken()
  const newHash = await hashToken(newRefreshToken)

  const updated = await db
    .update(tables.userSessions)
    .set({
      refreshTokenHash: newHash,
      previousRefreshTokenHash: oldHash,
      lastRefreshedAt: now,
      expiresAt: new Date(now.getTime() + SESSION_TTL_MS),
    })
    .where(eq(tables.userSessions.refreshTokenHash, oldHash))
    .returning({ id: tables.userSessions.id })

  const accessToken = await mintAccessToken(user, jwtSecrets)

  if (updated.length === 0) {
    return { kind: 'race_lost', user, accessToken }
  }

  return { kind: 'rotated', user, accessToken, refreshToken: newRefreshToken }
}

export async function bumpSessionVersion<S extends Record<string, unknown>>(
  db: Querier<S>,
  usersTable: typeof authUsers,
  userId: string,
): Promise<number> {
  const [row] = await db
    .update(usersTable)
    .set({ sessionVersion: sql`${usersTable.sessionVersion} + 1` })
    .where(eq(usersTable.id, userId))
    .returning({ sessionVersion: usersTable.sessionVersion })
  if (!row) throw new Error(`bumpSessionVersion: user not found (${userId})`)
  return row.sessionVersion
}
