"""Task-owned transparent OpenAI-compatible alias and one-forward-per-case recorder.

The native CE client remains unmodified. Only model-identifying metadata is mapped.
Generated choices/content are preserved. Private upstream and served bytes are retained.
This program does not start/load/download a model or issue generation on startup.
"""
import argparse
import datetime
import hashlib
import http.server
import json
import os
import pathlib
import secrets
import threading
import urllib.error
import urllib.parse
import urllib.request

ROOT = pathlib.Path(__file__).resolve().parent
RECORDS = ROOT / 'records' / 'alias-bridge-v1'
RECORDS.mkdir(exist_ok=True)
PRIVATE = RECORDS / 'private'
PUBLIC = RECORDS / 'public'
PRIVATE.mkdir(exist_ok=True)
PUBLIC.mkdir(exist_ok=True)
ALIAS = 'uagentkit-qwen2.5-coder-1.5b-q4km'
UPSTREAM = 'http://127.0.0.1:36280'
ALLOWED_ORIGIN = 'http://127.0.0.1:37416'
UPSTREAM_MODEL = json.loads((ROOT / 'records' / 'existing-provider-models.body').read_text(encoding='utf-8'))['data'][0]['id']
CASES = json.loads((ROOT / 'records' / 'proposed-cases-v1.json').read_text(encoding='utf-8'))
CASE_INPUTS = {row['id']: (ROOT / 'records' / row['promptFile']).read_text(encoding='utf-8') for row in CASES}
LOCK = threading.RLock()
STATE_PATH = PRIVATE / 'case-state.json'
CONTROL_TOKEN_PATH = PRIVATE / 'control-token.txt'
if CONTROL_TOKEN_PATH.exists():
    CONTROL_TOKEN = CONTROL_TOKEN_PATH.read_text(encoding='utf-8').strip()
else:
    CONTROL_TOKEN = secrets.token_urlsafe(32)
    CONTROL_TOKEN_PATH.write_text(CONTROL_TOKEN, encoding='utf-8')

def utc():
    return datetime.datetime.now(datetime.timezone.utc).isoformat()

def digest(data):
    return hashlib.sha256(data).hexdigest()

def json_bytes(value):
    return (json.dumps(value, ensure_ascii=False, separators=(',', ':')) + '\n').encode('utf-8')

def write_json(path, value):
    temp = path.with_suffix(path.suffix + '.tmp')
    temp.write_text(json.dumps(value, ensure_ascii=False, indent=2) + '\n', encoding='utf-8')
    os.replace(temp, path)

def normalize_input(text):
    # The original UTF-8 prompt bytes remain frozen/copied and hashes are retained.
    # Allow only the declared native selected-block Markdown serialization differences.
    # Retain every original/native byte and all generated choices/content unchanged.
    # Native selected-block Markdown serialization may add blank lines or normalize
    # a list marker. Preserve every nonblank line's words, punctuation and order.
    import re
    lines = []
    for line in text.replace('\r\n', '\n').split('\n'):
        if not line.strip():
            continue
        line = line.strip()
        line = re.sub(r'^[*-] ', '- ', line)
        lines.append(line)
    return '\n'.join(lines)

if STATE_PATH.exists():
    STATE = json.loads(STATE_PATH.read_text(encoding='utf-8'))
    if set(STATE['cases']) != set(CASE_INPUTS):
        raise RuntimeError('Existing cumulative case state does not match proposed cases')
else:
    STATE = {
        'createdAt': utc(), 'sequence': 0, 'armedCase': None,
        'cases': {case_id: {'generationAttempts': 0, 'backendForwards': 0, 'completed': False, 'forwardReserved': False} for case_id in CASE_INPUTS},
        'counters': {'requests': 0, 'metadataRequests': 0, 'metadataBackendForwards': 0, 'controlRequests': 0, 'preflightRequests': 0, 'generationAttempts': 0, 'generationBackendForwards': 0, 'generationDenied': 0},
    }
    write_json(STATE_PATH, STATE)

def save_state():
    write_json(STATE_PATH, STATE)

def append_event(event):
    # Only a small projection is public: never input/output/header/path bodies.
    with LOCK:
        with (PRIVATE / 'events.jsonl').open('a', encoding='utf-8') as f:
            f.write(json.dumps(event, ensure_ascii=False, separators=(',', ':')) + '\n')
        public = {key: event[key] for key in ['sequence', 'at', 'kind', 'method', 'path', 'caseId', 'action', 'status', 'backendForwarded', 'modelAlias', 'requestBytes', 'requestSha256', 'upstreamRequestBytes', 'upstreamRequestSha256', 'upstreamResponseBytes', 'upstreamResponseSha256', 'servedResponseBytes', 'servedResponseSha256', 'inputEquivalent', 'choiceContentUnchanged', 'errorCode'] if key in event}
        with (PUBLIC / 'events.jsonl').open('a', encoding='utf-8') as f:
            f.write(json.dumps(public, ensure_ascii=False, separators=(',', ':')) + '\n')

def state_projection():
    with LOCK:
        return json.loads(json.dumps({'recordedAt': utc(), 'modelAlias': ALIAS, 'armedCase': STATE['armedCase'], 'cases': STATE['cases'], 'counters': STATE['counters']}))

def finish_event(event):
    append_event(event)
    with LOCK:
        save_state()
        write_json(PUBLIC / 'state.json', state_projection())

def map_metadata(value, in_choices=False):
    """Map alias metadata recursively; never touch the generated choices subtree."""
    if isinstance(value, dict):
        return {key: item if key == 'choices' else map_metadata(item, in_choices) for key, item in value.items()}
    if isinstance(value, list):
        return [map_metadata(item, in_choices) for item in value]
    if isinstance(value, str):
        return value.replace(UPSTREAM_MODEL, ALIAS)
    return value

def choices_preserved(original, mapped):
    if isinstance(original, dict) and isinstance(mapped, dict) and 'choices' in original:
        return original['choices'] == mapped.get('choices')
    return True

def last_user_text(payload):
    users = [msg for msg in payload.get('messages', []) if isinstance(msg, dict) and msg.get('role') == 'user']
    if not users:
        return None
    content = users[-1].get('content')
    if isinstance(content, str):
        return content
    if isinstance(content, list) and all(isinstance(part, dict) and part.get('type') == 'text' and isinstance(part.get('text'), str) for part in content):
        return ''.join(part['text'] for part in content)
    return None

class Handler(http.server.BaseHTTPRequestHandler):
    protocol_version = 'HTTP/1.1'
    server_version = 'uAgentKitAliasRecorder/1'

    def log_message(self, fmt, *args):
        print(utc(), fmt % args, flush=True)

    def origin_ok(self):
        return self.headers.get('Origin') in (None, ALLOWED_ORIGIN)

    def read_body(self):
        length = int(self.headers.get('Content-Length', '0'))
        if length > 1024 * 1024:
            raise ValueError('request_too_large')
        return self.rfile.read(length)

    def new_event(self, kind, body=b''):
        with LOCK:
            STATE['sequence'] += 1
            STATE['counters']['requests'] += 1
            seq = STATE['sequence']
            case_id = STATE['armedCase'] if kind == 'generation' else None
            event = {'sequence': seq, 'at': utc(), 'kind': kind, 'method': self.command, 'path': urllib.parse.urlsplit(self.path).path, 'caseId': case_id, 'modelAlias': ALIAS, 'requestBytes': len(body), 'requestSha256': digest(body), 'backendForwarded': False, 'origin': self.headers.get('Origin'), 'authorizationHeaderPresent': bool(self.headers.get('Authorization'))}
            if body and kind != 'control':
                path = PRIVATE / f'{seq:04d}-native-request.bin'
                path.write_bytes(body)
                event['nativeRequestFile'] = path.name
            save_state()
            return event

    def send_body(self, status, body, content_type='application/json', extra_headers=None):
        self.send_response(status)
        self.send_header('Content-Type', content_type)
        self.send_header('Content-Length', str(len(body)))
        self.send_header('Cache-Control', 'no-store')
        if self.headers.get('Origin') == ALLOWED_ORIGIN:
            self.send_header('Access-Control-Allow-Origin', ALLOWED_ORIGIN)
            self.send_header('Vary', 'Origin')
        for key, value in (extra_headers or {}).items():
            self.send_header(key, value)
        self.end_headers()
        self.wfile.write(body)
        self.wfile.flush()

    def reject(self, event, status, code, message):
        body = json_bytes({'error': {'message': message, 'type': 'task_owned_recording_guard', 'code': code}})
        event.update(status=status, action='denied', errorCode=code, servedResponseBytes=len(body), servedResponseSha256=digest(body))
        if event['kind'] == 'generation':
            with LOCK:
                STATE['counters']['generationDenied'] += 1
        finish_event(event)
        self.send_body(status, body)

    def do_OPTIONS(self):
        event = self.new_event('preflight')
        with LOCK:
            STATE['counters']['preflightRequests'] += 1
        if not self.origin_ok():
            return self.reject(event, 403, 'origin_not_allowed', 'Only the task-owned native web origin is allowed.')
        self.send_body(204, b'', extra_headers={'Access-Control-Allow-Methods': 'GET, POST, OPTIONS', 'Access-Control-Allow-Headers': 'Content-Type, Authorization', 'Access-Control-Max-Age': '0'})
        event.update(status=204, action='cors_preflight')
        finish_event(event)

    def control_auth(self):
        return secrets.compare_digest(self.headers.get('X-UAgentKit-Control', ''), CONTROL_TOKEN)

    def do_GET(self):
        path = urllib.parse.urlsplit(self.path).path
        if path == '/_control/status':
            event = self.new_event('control')
            with LOCK:
                STATE['counters']['controlRequests'] += 1
            if not self.control_auth():
                return self.reject(event, 403, 'control_token_required', 'Private task control authentication is required.')
            body = json_bytes(state_projection())
            self.send_body(200, body)
            event.update(status=200, action='read_state')
            return finish_event(event)
        event = self.new_event('metadata')
        with LOCK:
            STATE['counters']['metadataRequests'] += 1
        if not self.origin_ok():
            return self.reject(event, 403, 'origin_not_allowed', 'Only the task-owned native web origin is allowed.')
        if path in ('/health', '/_bridge/health'):
            body = json_bytes({'status': 'ok', 'scope': 'bridge-ready-only', 'modelAlias': ALIAS, 'armedCase': STATE['armedCase'], 'generationBackendForwards': STATE['counters']['generationBackendForwards']})
            self.send_body(200, body)
            event.update(status=200, action='bridge_health', servedResponseBytes=len(body), servedResponseSha256=digest(body))
            return finish_event(event)
        if path in ('/v1/models', '/models'):
            return self.forward(event, 'GET', '/v1/models', None)
        return self.reject(event, 404, 'unsupported_metadata_path', 'The recording bridge supports health and OpenAI models metadata only.')

    def do_POST(self):
        try:
            body = self.read_body()
        except (ValueError, OSError):
            event = self.new_event('invalid')
            return self.reject(event, 400, 'invalid_request_body', 'The request body is invalid or too large.')
        path = urllib.parse.urlsplit(self.path).path
        if path == '/_control/arm':
            event = self.new_event('control')
            with LOCK:
                STATE['counters']['controlRequests'] += 1
            if not self.control_auth():
                return self.reject(event, 403, 'control_token_required', 'Private task control authentication is required.')
            try:
                data = json.loads(body)
                case_id = data['caseId']
            except (ValueError, KeyError, TypeError):
                return self.reject(event, 400, 'invalid_case', 'A named original case is required.')
            with LOCK:
                if case_id not in STATE['cases']:
                    return self.reject(event, 400, 'unknown_case', 'This case is outside the original two-case plan.')
                case_state = STATE['cases'][case_id]
                if case_state['backendForwards'] or case_state['forwardReserved'] or case_state['completed']:
                    return self.reject(event, 409, 'case_budget_consumed', 'This original case cannot be re-armed after its single backend forward was reserved.')
                if STATE['armedCase'] not in (None, case_id):
                    return self.reject(event, 409, 'different_case_still_armed', 'Finish the already armed case first.')
                STATE['armedCase'] = case_id
                event.update(caseId=case_id, status=200, action='arm_case')
                finish_event(event)
                self.send_body(200, json_bytes({'armedCase': case_id, 'modelAlias': ALIAS, 'remainingBackendForwards': 1}))
                return
        event = self.new_event('generation', body)
        with LOCK:
            STATE['counters']['generationAttempts'] += 1
            if event['caseId']:
                STATE['cases'][event['caseId']]['generationAttempts'] += 1
        if not self.origin_ok():
            return self.reject(event, 403, 'origin_not_allowed', 'Only the task-owned native web origin is allowed.')
        if path not in ('/v1/chat/completions', '/chat/completions'):
            return self.reject(event, 404, 'unsupported_generation_path', 'Use the native OpenAI-compatible chat/completions route.')
        try:
            payload = json.loads(body)
            if not isinstance(payload, dict):
                raise ValueError()
        except ValueError:
            return self.reject(event, 400, 'invalid_generation_json', 'Native generation body must be a JSON object.')
        actual_input = last_user_text(payload)
        # A native retry after the original request is disarmed remains attributed
        # to its exact original prompt. It is recorded and denied, never forwarded.
        if event['caseId'] is None and isinstance(actual_input, str):
            matched = [case_id for case_id, expected in CASE_INPUTS.items() if normalize_input(actual_input) == normalize_input(expected)]
            if len(matched) == 1:
                with LOCK:
                    event['caseId'] = matched[0]
                    STATE['cases'][matched[0]]['generationAttempts'] += 1
        with LOCK:
            case_id = STATE['armedCase']
            if not case_id:
                consumed = event['caseId'] and STATE['cases'][event['caseId']]['backendForwards']
                code = 'case_budget_consumed' if consumed else 'case_not_armed'
                return self.reject(event, 409, code, 'Generation is disabled unless an original unconsumed case is explicitly armed.')
            event['caseId'] = case_id
            case_state = STATE['cases'][case_id]
            if case_state['backendForwards'] or case_state['forwardReserved']:
                return self.reject(event, 409, 'case_budget_consumed', 'The original case has already consumed its one backend forward.')
            if payload.get('model') != ALIAS:
                return self.reject(event, 400, 'wrong_model_alias', 'The request must use the fixed public model alias.')
            equivalent = isinstance(actual_input, str) and normalize_input(actual_input) == normalize_input(CASE_INPUTS[case_id])
            event.update(inputEquivalent=equivalent, expectedPromptSha256=digest(CASE_INPUTS[case_id].encode('utf-8')), actualLastUserSha256=digest(actual_input.encode('utf-8')) if isinstance(actual_input, str) else None)
            if not equivalent:
                return self.reject(event, 409, 'original_case_input_mismatch', 'Only the original armed case input may reach the real backend; ancillary title/suggestion/capability prompts are denied.')
            systems = [m.get('content') for m in payload.get('messages', []) if m.get('role') == 'system']
            expected_system = (ROOT / 'records' / 'expected-system-v1.txt').read_text(encoding='utf-8')
            if systems != [expected_system] or len(payload.get('messages', [])) != 2:
                return self.reject(event, 409, 'native_editor_system_mismatch', 'Only the frozen native inline editor system plus original note action may be forwarded.')
            if payload.get('tools') or payload.get('functions'):
                return self.reject(event, 409, 'tools_outside_text_case', 'The frozen proposed cases are text-only; Use the native inline editor; agent, tools and embedding are outside this case.')
            # Persist consumed budget before any network call. Timeout/error does not refund.
            case_state['forwardReserved'] = True
            case_state['backendForwards'] += 1
            STATE['counters']['generationBackendForwards'] += 1
            STATE['armedCase'] = None
            save_state()
        upstream_payload = dict(payload)
        upstream_payload['model'] = UPSTREAM_MODEL
        upstream_body = json_bytes(upstream_payload)
        private_path = PRIVATE / f"{event['sequence']:04d}-upstream-request.bin"
        private_path.write_bytes(upstream_body)
        event.update(upstreamRequestFile=private_path.name, upstreamRequestBytes=len(upstream_body), upstreamRequestSha256=digest(upstream_body), nativeModelAlias=ALIAS, realBackendModelId=UPSTREAM_MODEL, bodyMapping='model field only; native prompt/settings/tools unchanged')
        self.forward(event, 'POST', '/v1/chat/completions', upstream_body)

    def forward(self, event, method, upstream_path, data):
        seq = event['sequence']
        raw_path = PRIVATE / f'{seq:04d}-upstream-response.bin'
        served_path = PRIVATE / f'{seq:04d}-served-response.bin'
        event['backendForwarded'] = True
        event['upstreamPath'] = upstream_path
        if event['kind'] == 'metadata':
            with LOCK:
                STATE['counters']['metadataBackendForwards'] += 1
        request = urllib.request.Request(UPSTREAM + upstream_path, data=data, method=method, headers={'Content-Type': 'application/json', 'Accept': 'text/event-stream, application/json', 'Accept-Encoding': 'identity'})
        raw_hash, served_hash = hashlib.sha256(), hashlib.sha256()
        raw_count = served_count = 0
        all_choices_unchanged = True
        try:
            try:
                response = urllib.request.urlopen(request, timeout=180)
            except urllib.error.HTTPError as exc:
                response = exc
            status = response.status if hasattr(response, 'status') else response.code
            content_type = response.headers.get('Content-Type', 'application/json')
            event.update(status=status, action='real_backend_forward', upstreamResponseFile=raw_path.name, servedResponseFile=served_path.name)
            with response, raw_path.open('wb') as raw_file, served_path.open('wb') as served_file:
                if 'text/event-stream' in content_type:
                    self.send_response(status)
                    self.send_header('Content-Type', content_type)
                    self.send_header('Cache-Control', 'no-store')
                    self.send_header('Connection', 'close')
                    if self.headers.get('Origin') == ALLOWED_ORIGIN:
                        self.send_header('Access-Control-Allow-Origin', ALLOWED_ORIGIN)
                        self.send_header('Vary', 'Origin')
                    self.end_headers()
                    self.close_connection = True
                    while True:
                        line = response.readline()
                        if not line:
                            break
                        raw_file.write(line)
                        raw_hash.update(line)
                        raw_count += len(line)
                        served = line
                        if line.startswith(b'data:') and line[5:].strip() != b'[DONE]':
                            try:
                                original = json.loads(line[5:].strip())
                                mapped = map_metadata(original)
                                all_choices_unchanged = all_choices_unchanged and choices_preserved(original, mapped)
                                served = b'data: ' + json.dumps(mapped, ensure_ascii=False, separators=(',', ':')).encode('utf-8') + b'\n'
                            except (ValueError, UnicodeError):
                                event['unparsedSseLineForwardedUnchanged'] = True
                        served_file.write(served)
                        served_hash.update(served)
                        served_count += len(served)
                        self.wfile.write(served)
                        self.wfile.flush()
                else:
                    raw = response.read()
                    raw_file.write(raw)
                    raw_hash.update(raw)
                    raw_count = len(raw)
                    served = raw
                    try:
                        original = json.loads(raw)
                        mapped = map_metadata(original)
                        all_choices_unchanged = choices_preserved(original, mapped)
                        served = json_bytes(mapped)
                    except (ValueError, UnicodeError):
                        event['unparsedBodyForwardedUnchanged'] = True
                    served_file.write(served)
                    served_hash.update(served)
                    served_count = len(served)
                    self.send_body(status, served, content_type)
            event.update(upstreamResponseBytes=raw_count, upstreamResponseSha256=raw_hash.hexdigest(), servedResponseBytes=served_count, servedResponseSha256=served_hash.hexdigest(), choiceContentUnchanged=all_choices_unchanged, transportComplete=True)
        except Exception as exc:
            # Preserve the failure, do not retry or refund a reserved original case.
            event.update(action='forward_failed_no_retry', errorCode=type(exc).__name__, privateError=str(exc), upstreamResponseBytes=raw_count, upstreamResponseSha256=raw_hash.hexdigest(), servedResponseBytes=served_count, servedResponseSha256=served_hash.hexdigest(), choiceContentUnchanged=all_choices_unchanged, transportComplete=False)
            if not served_count:
                try:
                    self.send_body(502, json_bytes({'error': {'message': 'The real backend request failed; this recording bridge does not retry or refund the case budget.', 'type': 'task_owned_transport_error', 'code': 'upstream_failed'}}))
                except Exception:
                    pass
        finally:
            if event['kind'] == 'generation' and event.get('caseId'):
                with LOCK:
                    STATE['cases'][event['caseId']]['completed'] = True
            finish_event(event)

def main():
    parser = argparse.ArgumentParser()
    parser.add_argument('--port', type=int, required=True)
    args = parser.parse_args()
    source_bytes = pathlib.Path(__file__).read_bytes()
    configuration = {
        'createdAt': utc(), 'host': '127.0.0.1', 'port': args.port,
        'endpoint': f'http://127.0.0.1:{args.port}/v1', 'modelAlias': ALIAS,
        'realBackend': UPSTREAM, 'realBackendModelId': UPSTREAM_MODEL,
        'allowedOrigin': ALLOWED_ORIGIN, 'controlTokenFile': str(CONTROL_TOKEN_PATH),
        'bridgeSourceSha256': digest(source_bytes), 'stateFile': str(STATE_PATH),
        'maxBackendForwardsPerCase': 1, 'cumulativeCaseBudgetPersistsAcrossRestart': True,
        'inputEquivalence': 'Transport CRLF, blank-line spacing, leading/trailing line whitespace and * versus - list marker equivalence only; all nonblank line wording/punctuation/order preserved. Raw native/request/original bytes retained.',
        'inputExpectedHashes': {case_id: digest(text.encode('utf-8')) for case_id, text in CASE_INPUTS.items()},
        'modelMapping': 'Native alias request model -> exact real backend ID; output metadata identifiers -> alias. Choices subtree unchanged; raw upstream and served response bytes retained privately.',
        'startupInferenceRequests': 0,
    }
    write_json(PRIVATE / 'configuration.json', configuration)
    public = {key: value for key, value in configuration.items() if key not in ('realBackendModelId', 'controlTokenFile', 'stateFile')}
    write_json(PUBLIC / 'configuration.json', public)
    write_json(PUBLIC / 'state.json', state_projection())
    server = http.server.ThreadingHTTPServer(('127.0.0.1', args.port), Handler)
    print(json.dumps({'startedAt': utc(), 'pid': os.getpid(), 'endpoint': configuration['endpoint'], 'modelAlias': ALIAS, 'maxBackendForwardsPerCase': 1, 'armedCase': STATE['armedCase'], 'startupInferenceRequests': 0}), flush=True)
    server.serve_forever(poll_interval=0.5)

if __name__ == '__main__':
    main()
