"""
Volt, Solara's own AI model, as a cPanel "Setup Python App" (Passenger / WSGI).

It answers the same three calls as llama.cpp's llama-server, so the plugin needs no changes:
  GET  /health               200 when the model is loaded, 503 while it loads (no key needed)
  GET  /v1/models            the model name (key needed)
  POST /v1/chat/completions  OpenAI-style chat request with response_format json_schema (key needed)

Settings (cPanel > Setup Python App > Environment variables), all optional:
  SOLARA_AI_KEY    access key (required; check.py suggests one)
  SOLARA_MODEL     model file (default models/model.gguf next to this file)
  SOLARA_THREADS   CPU threads (default 2; shared hosting usually caps you at 1-2 cores)
  SOLARA_CTX       context size in tokens (default 3072; Solara's prompts are under 1,600)
  SOLARA_SLOTS     1 or 2 copies of the conversation cache (default 1; 2 is faster but needs ~350 MB more)
"""
import hmac
import json
import os
import secrets
import sys
import threading
import time

HERE = os.path.dirname(os.path.abspath(__file__))
MODEL_PATH = os.environ.get('SOLARA_MODEL') or os.path.join(HERE, 'models', 'model.gguf')
THREADS = max(1, int(os.environ.get('SOLARA_THREADS') or 2))
CTX = max(2048, int(os.environ.get('SOLARA_CTX') or 3072))
SLOTS = 2 if os.environ.get('SOLARA_SLOTS') == '2' else 1
MODEL_NAME = 'solara'


def access_key():
    """From the app's environment variables only: never a file, because the app folder may be inside the web root."""
    return (os.environ.get('SOLARA_AI_KEY') or '').strip()


KEY = access_key()

# --------------------------------------------------------------------------------------------------
# The model. Loaded in the background so Passenger starts the app at once and /health can say "loading".
# --------------------------------------------------------------------------------------------------

_slots = []
_ready = threading.Event()
_load_error = []
_lock = threading.Lock()

try:
    import fcntl  # one generation at a time across Passenger processes, so they don't fight over the CPU
except ImportError:
    fcntl = None


def _new_llama():
    from llama_cpp import Llama
    return Llama(model_path=MODEL_PATH, n_ctx=CTX, n_threads=THREADS, n_threads_batch=THREADS,
                 n_batch=512, use_mmap=True, verbose=False)


def _load():
    try:
        if not os.path.exists(MODEL_PATH):
            raise RuntimeError('Model file not found: ' + MODEL_PATH)
        _slots.append(_new_llama())
    except Exception as e:  # noqa: BLE001 - reported on /health
        _load_error.append(str(e))
    finally:
        _ready.set()


threading.Thread(target=_load, daemon=True).start()


def _common_prefix(a, b):
    n = 0
    for x, y in zip(a, b):
        if x != y:
            break
        n += 1
    return n


def _pick_slot(tokens):
    """The copy whose cached conversation shares the most of this prompt (llama-server's slots, in miniature)."""
    best, best_len = _slots[0], -1
    for llm in _slots:
        n = _common_prefix(list(getattr(llm, '_input_ids', [])), tokens)
        if n > best_len:
            best, best_len = llm, n
    if SLOTS == 2 and len(_slots) == 1 and best_len < len(tokens) // 2:
        _slots.append(_new_llama())  # weights are memory-mapped, so the second copy shares them
        return _slots[1]
    return best


def chat_prompt(messages):
    """Qwen3's chat template with thinking switched off (what llama-server does with enable_thinking=false)."""
    out = []
    for m in messages:
        role = m.get('role', 'user')
        content = m.get('content') or ''
        if isinstance(content, list):
            content = ''.join(p.get('text', '') for p in content if isinstance(p, dict))
        out.append('<|im_start|>' + role + '\n' + content.strip() + '<|im_end|>\n')
    out.append('<|im_start|>assistant\n<think>\n\n</think>\n\n')
    return ''.join(out)


def generate(messages, schema, max_tokens):
    from llama_cpp import LlamaGrammar
    prompt = chat_prompt(messages)
    grammar = LlamaGrammar.from_json_schema(json.dumps(schema), verbose=False) if schema else None
    lock_file = None
    with _lock:
        try:
            if fcntl:
                lock_file = open(os.path.join(HERE, '.generate.lock'), 'w')
                fcntl.flock(lock_file, fcntl.LOCK_EX)
            tokens = _slots[0].tokenize(prompt.encode('utf-8'), add_bos=False, special=True)
            if len(tokens) + max_tokens > CTX:
                raise ValueError('Prompt is %d tokens; raise SOLARA_CTX above %d.' % (len(tokens), len(tokens) + max_tokens))
            llm = _pick_slot(tokens)
            res = llm.create_completion(prompt=prompt, max_tokens=max_tokens, temperature=0, grammar=grammar,
                                        stop=['<|im_end|>'])
        finally:
            if lock_file:
                fcntl.flock(lock_file, fcntl.LOCK_UN)
                lock_file.close()
    choice = res['choices'][0]
    return choice.get('text', ''), choice.get('finish_reason') or 'stop', res.get('usage', {})


# --------------------------------------------------------------------------------------------------
# HTTP
# --------------------------------------------------------------------------------------------------

def _json(start_response, status, body):
    data = json.dumps(body).encode('utf-8')
    start_response(status, [('Content-Type', 'application/json'), ('Content-Length', str(len(data))),
                            ('Cache-Control', 'no-store')])
    return [data]


def _error(start_response, status, message):
    return _json(start_response, status, {'error': {'message': message}})


def _authorised(environ):
    sent = environ.get('HTTP_X_SOLARA_KEY', '')  # some Apache set-ups drop the Authorization header
    if not sent:
        auth = environ.get('HTTP_AUTHORIZATION') or environ.get('REDIRECT_HTTP_AUTHORIZATION') or ''
        sent = auth[7:] if auth.lower().startswith('bearer ') else ''
    return len(KEY) >= 20 and bool(sent) and hmac.compare_digest(sent.strip().encode(), KEY.encode())


def _path(environ):
    path = environ.get('PATH_INFO') or '/'
    for route in ('/health', '/v1/models', '/v1/chat/completions'):
        if path.endswith(route):  # works whether the app sits on a subdomain or under a folder
            return route
    return path


def application(environ, start_response):
    method = environ.get('REQUEST_METHOD', 'GET')
    path = _path(environ)

    if path == '/health' or path == '/':
        if len(KEY) < 20:
            return _error(start_response, '500 Internal Server Error',
                          'Add the environment variable SOLARA_AI_KEY (20+ characters) and restart the app.')
        if _load_error:
            return _error(start_response, '500 Internal Server Error', _load_error[0])
        if not _ready.is_set():
            return _json(start_response, '503 Service Unavailable', {'status': 'loading model'})
        return _json(start_response, '200 OK', {'status': 'ok'})

    if not _authorised(environ):
        return _error(start_response, '401 Unauthorized', 'Invalid API Key')

    if path == '/v1/models' and method == 'GET':
        return _json(start_response, '200 OK', {'object': 'list', 'data': [
            {'id': MODEL_NAME, 'object': 'model', 'owned_by': 'solara', 'file': os.path.basename(MODEL_PATH)}]})

    if path == '/v1/chat/completions' and method == 'POST':
        try:
            length = int(environ.get('CONTENT_LENGTH') or 0)
            req = json.loads(environ['wsgi.input'].read(length).decode('utf-8') if length else '{}')
        except (ValueError, KeyError):
            return _error(start_response, '400 Bad Request', 'Body must be JSON')
        messages = req.get('messages') or []
        if not messages:
            return _error(start_response, '400 Bad Request', 'messages is required')
        rf = req.get('response_format') or {}
        schema = (rf.get('json_schema') or {}).get('schema') or rf.get('schema')
        max_tokens = min(int(req.get('max_tokens') or 256), 1024)

        if not _ready.wait(90):
            return _error(start_response, '503 Service Unavailable', 'Model is still loading')
        if _load_error:
            return _error(start_response, '500 Internal Server Error', _load_error[0])
        started = time.time()
        try:
            text, finish, usage = generate(messages, schema, max_tokens)
        except Exception as e:  # noqa: BLE001
            print('generate failed: %s' % e, file=sys.stderr)
            return _error(start_response, '500 Internal Server Error', str(e))
        return _json(start_response, '200 OK', {
            'id': 'chatcmpl-' + secrets.token_hex(8),
            'object': 'chat.completion',
            'created': int(started),
            'model': MODEL_NAME,
            'choices': [{'index': 0, 'finish_reason': finish,
                         'message': {'role': 'assistant', 'content': text}}],
            'usage': {'prompt_tokens': usage.get('prompt_tokens', 0),
                      'completion_tokens': usage.get('completion_tokens', 0),
                      'total_tokens': usage.get('total_tokens', 0)},
            'timings': {'seconds': round(time.time() - started, 2)},
        })

    return _error(start_response, '404 Not Found', 'Not found')
