const std = @import("std");

const zml = @import("zml");
const attention = zml.attention.attention;

const inference = @import("inference.zig");
const model = @import("model.zig");

pub const Session = struct {
    allocator: std.mem.Allocator,
    io: std.Io,
    platform: *const zml.Platform,
    model_buffers: *model.Buffers,
    compiled_model: *const inference.CompiledModel,
    config: *const model.Config,
    seqlen: u32,
    cache_buffers: zml.Bufferized(model.Cache),
    rng_buf: zml.Bufferized(zml.Tensor.Rng),
    generated_token_slice: zml.Slice,
    think_start: ?u32,
    think_end: ?u32,
    tokenizer: zml.tokenizer.Tokenizer,

    pub fn init(
        allocator: std.mem.Allocator,
        io: std.Io,
        platform: *const zml.Platform,
        tokenizer: zml.tokenizer.Tokenizer,
        compiled_model: *const inference.CompiledModel,
        model_buffers: *model.Buffers,
    ) !Session {
        const seed: u128 = @intCast(std.Io.Clock.now(.real, io).toNanoseconds());
        return .{
            .allocator = allocator,
            .io = io,
            .platform = platform,
            .model_buffers = model_buffers,
            .compiled_model = compiled_model,
            .tokenizer = tokenizer,
            .config = &compiled_model.loaded_model.parsed_config.value,
            .seqlen = compiled_model.params.seqlen,
            .cache_buffers = try compiled_model.params.cache.initBuffers(allocator, io, platform, compiled_model.params.shardings.model),
            .rng_buf = try zml.Tensor.Rng.initBuffer(io, platform, .replicated, seed),
            .generated_token_slice = try .alloc(allocator, zml.Shape.init(.{ .batch = 1, .seq = 1 }, .u32)),
            .think_start = tokenizer.tokenId("<think>") orelse unreachable,
            .think_end = tokenizer.tokenId("</think>") orelse unreachable,
        };
    }

    pub fn deinit(self: *Session) void {
        model.Cache.unloadBuffers(&self.cache_buffers);
        zml.Tensor.Rng.deinitBuffer(&self.rng_buf);
        self.generated_token_slice.free(self.allocator);
    }

    pub fn tokenizePrompt(self: *const Session, allocator: std.mem.Allocator, prompt: []const u8) ![]const u32 {
        var encoder = try self.tokenizer.encoder();
        defer encoder.deinit();

        const im_start = self.tokenizer.tokenId("<|im_start|>") orelse return error.NoSuchToken;
        const im_end = self.tokenizer.tokenId("<|im_end|>") orelse return error.NoSuchToken;
        const newline = self.tokenizer.tokenId("\\n") orelse return error.NoSuchToken;

        var tokens: std.ArrayList(u32) = try .initCapacity(allocator, prompt.len);
        try tokens.appendSlice(allocator, &.{ self.config.bos_token_id, im_start });
        const user_tokens = try encoder.encodeAlloc(allocator, "user\n");
        defer allocator.free(user_tokens);
        try tokens.appendSlice(allocator, user_tokens);
        const prompt_tokens = try encoder.encodeAlloc(allocator, prompt);
        defer allocator.free(prompt_tokens);
        try tokens.appendSlice(allocator, prompt_tokens);
        try tokens.appendSlice(allocator, &.{ im_end, newline, im_start });
        const assistant_tokens = try encoder.encodeAlloc(allocator, "assistant\n");
        defer allocator.free(assistant_tokens);
        try tokens.appendSlice(allocator, assistant_tokens);
        return tokens.toOwnedSlice(allocator);
    }

    pub fn tokenizeTurn(self: *const Session, allocator: std.mem.Allocator, prompt: []const u8) ![]const u32 {
        var encoder = try self.tokenizer.encoder();
        defer encoder.deinit();

        const im_start = self.tokenizer.tokenId("<|im_start|>") orelse return error.NoSuchToken;
        const im_end = self.tokenizer.tokenId("<|im_end|>") orelse return error.NoSuchToken;
        const newline = self.tokenizer.tokenId("\\n") orelse return error.NoSuchToken;

        var tokens: std.ArrayList(u32) = try .initCapacity(allocator, prompt.len);
        try tokens.appendSlice(allocator, &.{ im_end, newline, im_start });
        const user_tokens = try encoder.encodeAlloc(allocator, "user\n");
        defer allocator.free(user_tokens);
        try tokens.appendSlice(allocator, user_tokens);
        const prompt_tokens = try encoder.encodeAlloc(allocator, prompt);
        defer allocator.free(prompt_tokens);
        try tokens.appendSlice(allocator, prompt_tokens);
        try tokens.appendSlice(allocator, &.{ im_end, newline, im_start });
        const assistant_tokens = try encoder.encodeAlloc(allocator, "assistant\n");
        defer allocator.free(assistant_tokens);
        try tokens.appendSlice(allocator, assistant_tokens);
        return tokens.toOwnedSlice(allocator);
    }

    pub fn runPrefill(self: *Session, all_tokens: []const u32) !void {
        const tokens_slice: zml.Slice = try .alloc(self.allocator, .init(.{ .batch = 1, .seq = self.seqlen }, .u32));
        defer tokens_slice.free(self.allocator);
        const tokens = tokens_slice.items(u32);
        @memset(tokens, self.config.pad_token_id);
        @memcpy(tokens[0..all_tokens.len], all_tokens);

        var tokens_buf: zml.Buffer = try .fromSlice(self.io, self.platform, tokens_slice, .replicated);
        defer tokens_buf.deinit();

        const token_pos_slice: zml.Slice = .init(zml.Shape.init(.{ .batch = 1 }, .u32), std.mem.sliceAsBytes(&[_]u32{0}));
        var tokens_pos_buf: zml.Buffer = try .fromSlice(self.io, self.platform, token_pos_slice, .replicated);
        defer tokens_pos_buf.deinit();

        const actual_seq_len_slice: zml.Slice = .init(zml.Shape.init(.{}, .u32), std.mem.sliceAsBytes(&[_]u32{@intCast(all_tokens.len)}));
        var actual_seq_len_buf: zml.Buffer = try .fromSlice(self.io, self.platform, actual_seq_len_slice, .replicated);
        defer actual_seq_len_buf.deinit();

        const params = self.compiled_model.params;
        var attention_metadata_buffers: zml.Bufferized(attention.Metadata) = switch (params.attention_metadata) {
            .metal_fa => .{ .metal_fa = .{ .num_tokens = try .scalar(self.io, self.platform, all_tokens.len, .u32) } },
            else => try params.attention_metadata.initBuffer(self.io, self.platform, params.shardings.model),
        };
        defer attention.Metadata.deinitBuffer(&attention_metadata_buffers);

        try self.compiled_model.prefill.run(.{
            .allocator = self.allocator,
            .io = self.io,
            .platform = self.platform,
            .model_buffers = self.model_buffers,
            .tokens_buf = &tokens_buf,
            .tokens_pos_buf = &tokens_pos_buf,
            .actual_seq_len_buf = &actual_seq_len_buf,
            .rng_buf = &self.rng_buf,
            .cache_buffers = &self.cache_buffers,
            .attention_metadata_buffers = attention_metadata_buffers,
        });

        try tokens_buf.toSlice(self.io, tokens_slice);
        self.generated_token_slice.items(u32)[0] = tokens_slice.items(u32)[all_tokens.len - 1];
    }

    pub fn runDecode(self: *Session, all_tokens: *std.ArrayList(u32), stdout: *std.Io.Writer) !void {
        var decoder = try self.tokenizer.decoder();
        defer decoder.deinit();

        var current_token_buffer: zml.Buffer = try .fromSlice(self.io, self.platform, self.generated_token_slice, .replicated);
        defer current_token_buffer.deinit();

        const actual_seq_len_slice: zml.Slice = .init(zml.Shape.init(.{}, .u32), std.mem.sliceAsBytes(&[_]u32{0}));
        var actual_seq_len_buf: zml.Buffer = try .fromSlice(self.io, self.platform, actual_seq_len_slice, .replicated);
        defer actual_seq_len_buf.deinit();

        const out_tokens_buffer: []u8 = try self.allocator.alloc(u8, 1024);
        defer self.allocator.free(out_tokens_buffer);

        const params = self.compiled_model.params;
        var attention_metadata_buffers = try params.attention_metadata.initBuffer(self.io, self.platform, params.shardings.model);
        defer attention.Metadata.deinitBuffer(&attention_metadata_buffers);

        generation: while (true) {
            const token_id = self.generated_token_slice.items(u32)[0];

            if (token_id == self.config.eos_token_id) break :generation;

            const token = try decoder.feedOne(token_id, out_tokens_buffer);
            if (self.think_start) |think_start| if (token_id == think_start) {
                try stdout.writeAll("\x1b[2m");
            };
            try stdout.writeAll(token);
            if (self.think_end) |think_end| if (token_id == think_end) {
                try stdout.writeAll("\x1b[0m");
            };
            try stdout.flush();

            try all_tokens.append(self.allocator, token_id);
            if (all_tokens.items.len >= self.seqlen) break :generation;

            const token_pos_slice: zml.Slice = .init(zml.Shape.init(.{ .batch = 1 }, .u32), std.mem.sliceAsBytes(&[_]u32{@intCast(all_tokens.items.len)}));
            var token_pos_buffer: zml.Buffer = try .fromSlice(self.io, self.platform, token_pos_slice, .replicated);
            defer token_pos_buffer.deinit();

            try self.compiled_model.decode.run(.{
                .allocator = self.allocator,
                .io = self.io,
                .platform = self.platform,
                .model_buffers = self.model_buffers,
                .tokens_buf = &current_token_buffer,
                .tokens_pos_buf = &token_pos_buffer,
                .actual_seq_len_buf = &actual_seq_len_buf,
                .rng_buf = &self.rng_buf,
                .cache_buffers = &self.cache_buffers,
                .attention_metadata_buffers = attention_metadata_buffers,
            });

            try current_token_buffer.toSlice(self.io, self.generated_token_slice);
        }

        try stdout.writeAll(try decoder.finalize(out_tokens_buffer));
        try stdout.flush();
    }
};
