#!/usr/bin/env python3
"""Confer employee usage helper. Only normalized usage leaves this computer.

Provider credentials remain managed by the official clients. Python standard library.
Run `python3 employee-helper.py enroll` and then `python3 employee-helper.py watch`.
"""
import argparse
from contextlib import contextmanager
import datetime
import getpass
import hashlib
import json
import os
from pathlib import Path
import queue
import secrets
import subprocess
import sys
import threading
import time
import urllib.request
import urllib.error
from urllib.parse import urlparse

ROOT = Path.home() / '.confer-capacity'


def hashed(value):
    return hashlib.sha256(value.encode()).hexdigest()


def save(path, value):
    path.parent.mkdir(parents=True,exist_ok=True,mode=0o700)
    temp = path.with_name(path.name + '.' + secrets.token_hex(6))
    fd = os.open(temp,os.O_WRONLY | os.O_CREAT | os.O_EXCL,0o600)
    with os.fdopen(fd,'w') as f:
        json.dump(value,f)
        f.flush()
        os.fsync(f.fileno())
    os.replace(temp,path)


@contextmanager
def local_lock(path):
    path.parent.mkdir(parents=True,exist_ok=True,mode=0o700)
    with path.open('a+b') as f:
        if os.name == 'nt':
            import msvcrt
            f.seek(0);f.write(b'0');f.flush();f.seek(0)
            msvcrt.locking(f.fileno(),msvcrt.LK_LOCK,1)
        else:
            import fcntl
            fcntl.flock(f,fcntl.LOCK_EX)
        try:yield
        finally:
            if os.name == 'nt':
                f.seek(0);msvcrt.locking(f.fileno(),msvcrt.LK_UNLCK,1)
            else:fcntl.flock(f,fcntl.LOCK_UN)


class NoRedirect(urllib.request.HTTPRedirectHandler):
    def redirect_request(self, *args, **kwargs):
        return None


def post(config, endpoint, body, authenticated=True):
    headers={'Content-Type':'application/json','User-Agent':'ConferCapacity/1.0'}
    if authenticated:
        headers['Authorization']='Bearer '+config['secret']
    request=urllib.request.Request(config['url']+'/api/employees/'+endpoint,
        json.dumps(body).encode(), headers, method='POST')
    with urllib.request.build_opener(NoRedirect).open(request,timeout=15) as response:
        return json.load(response)


def native_codex():
    """Ask the installed official client to read account/limits; no inference turn."""
    p=subprocess.Popen(['codex','app-server'],stdin=subprocess.PIPE,stdout=subprocess.PIPE,
                       stderr=subprocess.DEVNULL,text=True)
    messages=queue.Queue()
    def read():
        for line in p.stdout:
            try: messages.put(json.loads(line))
            except ValueError: pass
        messages.put(None)
    reader=threading.Thread(target=read,daemon=True)
    reader.start()
    def rpc(n, method, params=None):
        p.stdin.write(json.dumps({'id':n,'method':method,'params':params})+'\n')
        p.stdin.flush()
        end=time.monotonic()+20
        while time.monotonic()<end:
            item=messages.get(timeout=max(.01,end-time.monotonic()))
            if item is None: raise RuntimeError('client exited')
            if item.get('id') == n:
                if 'error' in item: raise RuntimeError('client request unavailable')
                return item['result']
        raise RuntimeError('client timed out')
    try:
        rpc(1,'initialize',{'clientInfo':{'name':'confer_capacity','version':'1.0.0'}})
        p.stdin.write('{"method":"initialized"}\n'); p.stdin.flush()
        account=rpc(2,'account/read',{'refreshToken':False}).get('account') or {}
        limits=rpc(3,'account/rateLimits/read')
        account_id=account.get('email')
        buckets=limits.get('rateLimitsByLimitId') or {'codex':limits.get('rateLimits') or {}}
        windows=[]
        for key,bucket in buckets.items():
            for field in ('primary','secondary'):
                w=bucket.get(field)
                if not isinstance(w,dict): continue
                used,duration,reset=w.get('usedPercent'),w.get('windowDurationMins'),w.get('resetsAt')
                if any(v is None for v in (used,duration,reset)): continue
                windows.append({'bucket':key,'seconds':int(duration*60),'used_percent':used,'resets_at':reset})
        return {'provider':'codex','account':hashed(account_id.casefold()) if account_id else None,
                'status':'ok' if windows else 'unavailable','windows':windows,'observed_at':time.time()}
    finally:
        p.terminate()
        try: p.wait(timeout=3)
        except subprocess.TimeoutExpired: p.kill(); p.wait()
        p.stdin.close(); p.stdout.close()


def claude_statusline(raw, root=ROOT):
    """No network in the statusline: save only allowlisted native quota metadata."""
    windows=[]
    for name,seconds in [('five_hour',18000),('seven_day',604800)]:
        w=(raw.get('rate_limits') or {}).get(name) or {}
        if w.get('used_percentage') is not None and w.get('resets_at') is not None:
            windows.append({'bucket':'claude','seconds':seconds,'used_percent':w['used_percentage'],'resets_at':w['resets_at']})
    # A timer rerender is not a new provider observation. API-duration or quota changes are.
    fingerprint=hashed(json.dumps([windows, (raw.get('cost') or {}).get('total_api_duration_ms')],sort_keys=True))
    target=root/'claude-observation.json'
    session=hashed(str(raw.get('session_id','unknown')))
    session_file=root/'claude-sessions'/(session+'.json')
    with local_lock(root/'statusline.lock'):
        try: previous=json.loads(session_file.read_text())
        except (OSError,ValueError): previous={}
        if fingerprint != previous.get('fingerprint'):
            observation={'fingerprint':fingerprint,'provider':'claude','account':None,
                         'observed_at':time.time(),'status':'ok' if windows else 'unavailable','windows':windows}
            save(session_file,observation)
            save(target,observation)
    print('Confer usage: '+(' · '.join(f"{w['seconds']//3600}h {w['used_percent']}% used" for w in windows) or 'quota unavailable'))


def timestamp(value):
    try: return datetime.datetime.fromisoformat(value.replace('Z','+00:00')).timestamp()
    except (ValueError,TypeError,AttributeError): return None


def event_from_record(record, tool, session, progress=None):
    """Extract response counters only; never retain message or tool content."""
    at=timestamp(record.get('timestamp'))
    if at is None: return None
    if tool == 'claude':
        message=record.get('message') or {}
        usage=message.get('usage')
        if record.get('type') != 'assistant' or not usage or not message.get('id'): return None
        event_id=message['id']
        counts=[usage.get(k,0) for k in ('input_tokens','output_tokens','cache_read_input_tokens','cache_creation_input_tokens')]
    else:
        payload=record.get('payload') or {}
        if record.get('type') != 'event_msg' or payload.get('type') != 'token_count': return None
        info=payload.get('info') or {}
        total=info.get('total_token_usage')
        usage=info.get('last_token_usage')
        if not usage or not total or progress is None: return None
        values=[total.get(k,0) for k in ('input_tokens','output_tokens','cached_input_tokens')]
        if not all(type(n) is int and n >= 0 for n in values): return None
        prior=progress.get('total')
        if values == prior:return None  # Rate-limit-only events repeat previous token counts.
        epoch=progress.get('epoch',0)
        if prior and any(a<b for a,b in zip(values,prior)):
            epoch+=1;prior=None
        delta=[a-b for a,b in zip(values,prior)] if prior else [usage.get(k,0) for k in ('input_tokens','output_tokens','cached_input_tokens')]
        if not all(type(n) is int and n >= 0 for n in delta):return None
        progress.update(total=values,epoch=epoch)
        event_id=str(epoch)+':'+json.dumps(values)
        cached=min(delta[0],delta[2])
        counts=[delta[0]-cached,delta[1],cached,0]
    if not all(type(n) is int and 0 <= n <= 10**9 for n in counts): return None
    return dict(tool=tool,id=hashed(tool+':'+session+':'+event_id),at=at,
                **dict(zip(('input','output','cache_read','cache_write'),counts)))


def activity():
    """Rolling seven-day best-effort transcript coverage; scan capped at 64 MiB per file."""
    now=time.time()
    for tool,folder in [('claude',Path.home()/'.claude/projects'),('codex',Path(os.environ.get('CODEX_HOME',Path.home()/'.codex'))/'sessions')]:
        if not folder.exists(): continue
        for file in folder.rglob('*.jsonl'):
            if file.stat().st_mtime < now-7*86400 or file.stat().st_size > 64*1024*1024: continue
            with file.open(errors='replace') as lines:
                progress={}
                for line in lines:
                    try: record=json.loads(line)
                    except ValueError: continue
                    event=event_from_record(record,tool,file.stem,progress)
                    if event and now-7*86400 <= event['at'] <= now+60: yield event


def aggregate_activity(events):
    responses={}
    for event in events:
        key=event['id']
        old=responses.get(key)
        if old:
            event={**event,'at':min(old['at'],event['at']),
                   **{k:max(old[k],event[k]) for k in ('input','output','cache_read','cache_write')}}
        responses[key]=event
    return responses


def enroll(args):
    endpoint=args.url.rstrip('/')
    url=urlparse(endpoint)
    if url.scheme != 'https' or not url.hostname or url.path or url.username or url.password or url.query or url.fragment:
        raise ValueError('Enrollment requires an HTTPS origin')
    print('This connects this computer to Confer. It reports quota, reset times and optional token counts.\n'
          'Provider passwords, login tokens, prompts, file contents and paths are never uploaded.\n'
          'Remote execution is not enabled. An administrator must approve this device.')
    if input('Type CONNECT to continue: ').strip() != 'CONNECT': return
    config_path=ROOT/'config.json'
    if config_path.exists():
        config=json.loads(config_path.read_text())
        if config.get('device'): raise ValueError('Already enrolled. Revoke the device before removing its local configuration.')
        if config['url'] != endpoint: raise ValueError('Pending enrollment belongs to another server')
    else:
        config={'url':endpoint,'secret':secrets.token_urlsafe(32),'activity':args.activity,'label':args.label}
        save(config_path,config)  # Persist before redemption: a lost response is safely retryable.
    invite=getpass.getpass('Paste your one-time Confer invitation: ')
    result=post(config,'enroll',{'invite':invite,'device_secret':config['secret'],'label':config['label'],'consent':True},False)
    config.update(device=result['device'],employee=result['employee'])
    save(config_path,config)
    print('Enrolled for '+result['employee']+'. Status: '+result['status']+'. Ask your administrator to verify this device.')
    print('Use the official clients to sign in if needed: codex login; claude /login. Credentials stay there.')
    print('Configure the Claude statusLine command to run this script with statusline, preserving any existing status line.')


def collect(config):
    out=[]
    provider=config.get('provider')
    if provider in (None,'codex'):
        try: out.append(native_codex())
        except (OSError,RuntimeError,queue.Empty,ValueError,KeyError,TypeError):
            out.append({'provider':'codex','account':None,'status':'unavailable','windows':[],'observed_at':time.time()})
    if provider in (None,'claude'):
        try:
            claude=json.loads((ROOT/'claude-observation.json').read_text())
            claude.pop('fingerprint',None)
            if time.time()-claude['observed_at'] <= 600: out.append(claude)
        except (OSError,ValueError,KeyError): pass
    for o in out:
        # Stable for retries, monotonic with newly captured provider observations.
        o['seq']=int(o['observed_at']*1_000_000)
    post(config,'report',{'observations':out})
    if config.get('activity'):
        responses=aggregate_activity(e for e in activity() if provider is None or e['tool']==provider)
        checkpoint=ROOT/'activity-sent.json'
        try:sent=json.loads(checkpoint.read_text())
        except (OSError,ValueError):sent={}
        sent={key:value for key,value in sent.items() if key in responses}
        def flush(batch):
            post(config,'report',{'events':batch})
            for event in batch:sent[event['id']]=hashed(json.dumps(event,sort_keys=True))
            save(checkpoint,sent)  # Commit after each acknowledged batch, including partial polls.
        batch=[]
        for event in responses.values():
            if sent.get(event['id']) == hashed(json.dumps(event,sort_keys=True)):continue
            batch.append(event)
            if len(batch)==200:flush(batch);batch=[]
        if batch:flush(batch)


def main():
    os.umask(0o077)
    parser=argparse.ArgumentParser(description=__doc__)
    parser.add_argument('command',choices=['enroll','once','watch','statusline'])
    parser.add_argument('--url',default='https://auth.conferhub.top')
    parser.add_argument('--label',default='Work computer')
    parser.add_argument('--activity',action='store_true',help='Opt in to normalized local response token counts')
    args=parser.parse_args()
    if args.command == 'enroll': enroll(args); return
    if args.command == 'statusline': claude_statusline(json.load(sys.stdin)); return
    config=json.loads((ROOT/'config.json').read_text())
    while True:
        try:
            config=json.loads((ROOT/'config.json').read_text())
            collect(config)
            print('Usage report accepted.')
        except urllib.error.HTTPError as e: print('Report not accepted (HTTP '+str(e.code)+'). Check approval or enrollment.')
        except (OSError,ValueError): print('Report unavailable; will retry. Existing observations retain their original time.')
        if args.command == 'once': break
        time.sleep(300)


if __name__ == '__main__':
    try: main()
    except KeyboardInterrupt: pass
    except Exception: print('Unable to continue. Check enrollment configuration and installed clients.',file=sys.stderr); sys.exit(1)
