#ifndef RN_SLOT_MANAGER_H
#define RN_SLOT_MANAGER_H

#include "rn-slot.h"
#include "common.h"
#include "llama.h"
#include <vector>
#include <deque>
#include <map>
#include <functional>
#include <mutex>
#include <thread>
#include <atomic>
#include <condition_variable>

namespace rnllama {

// Forward declarations
struct llama_rn_context;
struct completion_token_output;

// Status structures for exposing slot manager state to JS
struct llama_rn_request_status {
    int32_t request_id;
    std::string type;           // "completion", "embedding", "rerank"
    std::string state;          // "queued", "processing_prompt", "generating", "done"
    size_t prompt_length;
    size_t tokens_generated;
    double prompt_ms;
    double generation_ms;
    double tokens_per_second;
};

struct llama_rn_parallel_status {
    int32_t n_parallel;
    int32_t active_slots;
    int32_t queued_requests;
    std::vector<llama_rn_request_status> requests;
};

enum class llama_rn_cancel_result {
    ACTIVE,
    QUEUED,
    NOT_FOUND,
};

// Queued request structure
struct llama_rn_queued_request {
    int32_t request_id;
    llama_rn_slot_task_type task_type;
    common_params params;
    std::vector<llama_token> prompt_tokens;
    std::function<void(const completion_token_output&)> on_token;
    std::function<void(llama_rn_slot*)> on_complete;

    // Media paths for multimodal
    std::vector<std::string> media_paths;
    std::string prompt_text;  // Original prompt text (needed for media processing)

    // Chat format parameters
    int chat_format;
    common_reasoning_format reasoning_format;
    std::string generation_prompt;
    std::string chat_parser;  // Serialized PEG parser for chat output parsing

    // Prefill text
    std::string prefill_text;

    // Embedding parameters
    int embd_normalize;
    std::function<void(int32_t, const std::vector<float>&)> on_embedding;

    // Rerank parameters
    std::vector<std::vector<llama_token>> rerank_prompt_tokens;
    std::function<void(int32_t, const std::vector<float>&)> on_rerank;

    // State management
    std::string load_state_path;       // File path to load state from before processing
    std::string save_state_path;       // File path to save state to after completion
    std::string save_prompt_state_path; // File path to save prompt state to after prompt processing
    int32_t load_state_size;           // Number of tokens to load (0 or -1 = all tokens)
    int32_t save_state_size;           // Number of tokens to save (0 or -1 = all tokens)

    llama_rn_queued_request() :
        request_id(-1),
        task_type(SLOT_TASK_TYPE_COMPLETION),
        chat_format(0),
        reasoning_format(COMMON_REASONING_FORMAT_NONE),
        embd_normalize(-1),
        load_state_size(-1),
        save_state_size(-1)
    {}
};

// Slot manager for parallel decoding
struct llama_rn_slot_manager {
    // Parent context reference
    llama_rn_context* parent_ctx;

    // Slot pool
    std::vector<llama_rn_slot> slots;
    int32_t n_parallel;                    // Number of parallel slots

    // Request queue
    std::deque<llama_rn_queued_request> queue_requests;

    // Request tracking
    std::map<int32_t, llama_rn_slot*> active_requests;  // request_id -> slot
    std::atomic<int32_t> next_request_id;

    // Batch processing
    llama_batch batch;
    int32_t n_batch;                       // Max batch size

    // Shared MTP speculative decoding state. llama.cpp's MTP driver is
    // multi-sequence, so queued slots borrow this instead of creating one
    // speculative context per slot.
    common_speculative *mtp_spec = nullptr;
    llama_context *mtp_spec_ctx = nullptr;
    int32_t mtp_spec_n_max = 0;
    int32_t mtp_spec_n_min = 0;
    float mtp_spec_p_min = 0.0f;
    bool mtp_spec_backend_sampling = true;
    int32_t mtp_spec_n_gpu_layers = -1;
    lm_ggml_type mtp_spec_cache_type_k = LM_GGML_TYPE_F16;
    lm_ggml_type mtp_spec_cache_type_v = LM_GGML_TYPE_F16;

    // Configuration
    float slot_prompt_similarity;          // Threshold for cache reuse (0.0-1.0)
    bool continuous_batching;              // Allow mixing prompt/generation

    // Processing loop control
    std::mutex slots_mutex;                // Mutex for thread-safe access to slots
    std::condition_variable slots_cv;      // Condition variable for efficient waiting
    std::thread processing_thread;         // Background processing thread
    std::atomic<bool> processing_active;   // Flag to control processing loop

    // Status subscription support
    std::map<int32_t, std::function<void(const llama_rn_parallel_status&)>> status_subscribers;
    std::mutex subscribers_mutex;
    int32_t next_subscriber_id = 1;

    // Constructor
    llama_rn_slot_manager(llama_rn_context* ctx);

    // Destructor
    ~llama_rn_slot_manager();

    // Initialization
    bool init(int32_t n_parallel, int32_t n_batch, int32_t n_ctx);

    // Request management
    int32_t reserve_request_id();

    int32_t queue_request(
        const common_params& params,
        const std::vector<llama_token>& prompt,
        const std::vector<std::string>& media_paths,
        const std::string& prompt_text,
        int chat_format,
        common_reasoning_format reasoning_format,
        const std::string& generation_prompt,
        const std::string& chat_parser,
        const std::string& prefill_text,
        const std::string& load_state_path,
        const std::string& save_state_path,
        const std::string& save_prompt_state_path,
        int32_t load_state_size,
        int32_t save_state_size,
        std::function<void(const completion_token_output&)> on_token,
        std::function<void(llama_rn_slot*)> on_complete,
        int32_t request_id = -1
    );

    int32_t queue_embedding_request(
        const std::vector<llama_token>& tokens,
        int embd_normalize,
        std::function<void(int32_t, const std::vector<float>&)> on_result,
        int32_t request_id = -1
    );

    int32_t queue_rerank_request(
        const std::string& query,
        const std::vector<std::string>& documents,
        int normalize,
        std::function<void(int32_t, const std::vector<float>&)> on_results,
        int32_t request_id = -1
    );

    // Slot management
    llama_rn_slot* get_available_slot(const std::vector<llama_token>& prompt);
    llama_rn_slot* get_slot_by_request_id(int32_t request_id);
    void release_slot(llama_rn_slot* slot);
    llama_rn_cancel_result cancel_request(int32_t request_id);

    // Processing loop management
    void start_processing_loop();
    void stop_processing_loop();

    // Main processing loop (protected by mutex)
    void update_slots();

    // Helper methods
    float compute_similarity(const std::vector<llama_token>& a,
                            const std::vector<llama_token>& b);
    void build_batch();
    bool process_batch();
    void sample_and_callback();

    // Finish a slot's generation: flush the UTF-8 gate, mark done, notify
    void complete_slot(llama_rn_slot & slot);
    common_speculative* ensure_mtp_speculative(common_params& params);
    llama_context* get_mtp_draft_context() const;
    void reset_mtp_speculative();

    // Process pending queue
    void process_pending_queue();

    // Release completed slots
    void release_completed_slots();

    // Status methods
    llama_rn_parallel_status get_status();
    bool has_pending_work();
    void notify_status_change();
    int32_t add_status_subscriber(std::function<void(const llama_rn_parallel_status&)> callback);
    void remove_status_subscriber(int32_t subscriber_id);
};

} // namespace rnllama

#endif /* RN_SLOT_MANAGER_H */
