/**
 * SSE usage extractor — parses SSE frames to capture token usage.
 *
 * When an AdapterRegistry adapter is provided it delegates to
 * `adapter.readUsageFromSse(eventType, dataJson)` so each provider can
 * define its own event-level token extraction.  Falls back to the
 * Anthropic-specific parser when no adapter is supplied.
 *
 * Returns a wrapping ReadableStream that passes all bytes through unchanged
 * while accumulating usage. When the stream ends (normally or on error),
 * onEnd() is called with `true` if any usage was captured, `false` otherwise.
 */
import { parseEvent, splitEvents } from "../codec/sseParser.ts";
import type { ProviderAdapter } from "../provider/ProviderAdapter.ts";

export interface SseUsageAcc {
  input_tokens: number;
  output_tokens: number;
  cache_read_input_tokens: number;
  cache_creation_input_tokens: number;
  /** Number of message_delta SSE frames observed (diagnostic). */
  deltas_seen: number;
  /** True if the wrapped stream was cancelled before completion (diagnostic). */
  was_cancelled: boolean;
}

/**
 * Parse a single SSE block (text between double-newlines) and update acc.
 * Returns true if the frame contributed any usage data.
 */
export function parseSseFrame(frame: string, acc: SseUsageAcc): boolean {
  const event = parseEvent(frame);
  const eventType = event.eventType ?? "";
  if (event.dataJson === null || typeof event.dataJson !== "object") return false;
  const payload = event.dataJson as Record<string, unknown>;

  if (eventType === "message_start") {
    const msg = (payload as { message?: { usage?: Record<string, number> } }).message;
    const usage = msg?.usage;
    if (!usage) return false;
    if (usage["input_tokens"] != null) acc.input_tokens = usage["input_tokens"];
    if (usage["cache_read_input_tokens"] != null) acc.cache_read_input_tokens = usage["cache_read_input_tokens"];
    if (usage["cache_creation_input_tokens"] != null) acc.cache_creation_input_tokens = usage["cache_creation_input_tokens"];
    // message_start also carries initial output_tokens (usually 1)
    if (usage["output_tokens"] != null) acc.output_tokens = Math.max(acc.output_tokens, usage["output_tokens"]);
    return true;
  }

  if (eventType === "message_delta") {
    const usage = (payload as { usage?: Record<string, number> }).usage;
    if (!usage) return false;
    if (usage["output_tokens"] != null) acc.output_tokens = Math.max(acc.output_tokens, usage["output_tokens"]);
    acc.deltas_seen++;
    return true;
  }

  return false;
}

/**
 * Dispatch a single SSE frame to accumulate usage.
 * When an adapter is provided, delegates to adapter.readUsageFromSse;
 * otherwise falls back to the legacy Anthropic-specific parser.
 */
function dispatchFrame(frame: string, acc: SseUsageAcc, adapter: ProviderAdapter | null | undefined): boolean {
  if (adapter) {
    const event = parseEvent(frame);
    if (event.isDone || event.dataJson === null) return false;
    const payload = event.dataJson;
    let eventType = event.eventType ?? "";
    if (!eventType) {
      const inferred = (payload as { type?: unknown } | null)?.type;
      if (typeof inferred === "string") eventType = inferred;
    }
    if (!eventType) return false;
    const block = adapter.readUsageFromSse(eventType, payload);
    if (!block) return false;
    if (block.input_tokens != null && (eventType !== "message_delta" || block.input_tokens > 0)) {
      acc.input_tokens = block.input_tokens;
    }
    if (block.output_tokens != null) acc.output_tokens = Math.max(acc.output_tokens ?? 0, block.output_tokens);
    if (block.cache_creation_input_tokens != null) acc.cache_creation_input_tokens = block.cache_creation_input_tokens;
    if (block.cache_read_input_tokens != null) acc.cache_read_input_tokens = block.cache_read_input_tokens;
    if (eventType === "message_delta") acc.deltas_seen++;
    return true;
  }
  return parseSseFrame(frame, acc);
}

export function wrapSseForUsage(
  body: ReadableStream<Uint8Array>,
  acc: SseUsageAcc,
  onEnd: (captured: boolean) => void,
  adapter?: ProviderAdapter | null,
  onError?: (err: unknown) => void,
): ReadableStream<Uint8Array> {
  const decoder = new TextDecoder("utf-8", { fatal: false });
  let buffer = "";
  let captured = false;
  let ended = false;

  const flushBuffer = () => {
    const { events, remainder } = splitEvents(buffer);
    buffer = remainder;
    for (const frame of events) {
      if (dispatchFrame(frame, acc, adapter)) captured = true;
    }
  };

  const finish = () => {
    if (ended) return;
    ended = true;
    // Flush any trailing content (upstream may not send final \n\n)
    if (buffer.trim()) {
      if (dispatchFrame(buffer.trim(), acc, adapter)) captured = true;
    }
    try {
      onEnd(captured);
    } catch (err) {
      onError?.(err);
    }
  };

  const reader = body.getReader();

  return new ReadableStream<Uint8Array>({
    async pull(controller) {
      try {
        const { done, value } = await reader.read();
        if (done) {
          finish();
          controller.close();
          return;
        }
        // Decode chunk (streaming mode handles split codepoints)
        buffer += decoder.decode(value, { stream: true });
        flushBuffer();
        controller.enqueue(value);
      } catch (err) {
        finish();
        onError?.(err);
        controller.error(err);
      }
    },
    cancel(reason) {
      acc.was_cancelled = true;
      finish();
      reader.cancel(reason);
    },
  });
}
