#!/usr/bin/env python3
"""iSynth external MCP client. Python >=3.10; pip install 'mcp==1.30.0'.

Credentials: ISYNTH_BASE_URL and ISYNTH_AUTOMATION_TOKEN in the environment.
The token is an iSynth scoped credential, NOT the Qwen/model API key.
See /guides/automation on your iSynth site for complete examples.
"""
import argparse
import asyncio
import hashlib
import io
import json
import os
import random
from datetime import datetime, timezone
from email.utils import parsedate_to_datetime
from pathlib import Path
from urllib.parse import urlsplit
from uuid import UUID
import zipfile
import httpx
from mcp import ClientSession
from mcp.client.streamable_http import streamablehttp_client


class RetryTransport(httpx.AsyncBaseTransport):
    """429 retries honor server cooldown; transient failures retry only safe calls."""
    def __init__(self, transport=None):
        self.transport=transport or httpx.AsyncHTTPTransport(trust_env=False)

    async def handle_async_request(self, request):
        body=await request.aread()
        try:
            payload=json.loads(body) if body else {}
        except ValueError:
            payload={}
        tool=payload.get('params',{}).get('name','')
        safe=request.method=='GET' or payload.get('method') in {'initialize','tools/list','notifications/initialized'} or tool.startswith('get_') or tool in {'start_upload','upload_batch','finish_upload','upload_records','retry_file_processing'}
        for attempt in range(6):
            response=None
            try:
                response=await self.transport.handle_async_request(httpx.Request(request.method,request.url,headers=request.headers,content=body,extensions=request.extensions))
                if attempt==5 or not (response.status_code==429 or safe and response.status_code in {408,500,502,503,504}):return response
            except httpx.TransportError:
                if not safe or attempt==5:raise
            delay=min(2**attempt,60)+random.uniform(0,1)
            if response is not None:
                value=response.headers.get('Retry-After','')
                try:wait=float(value)
                except ValueError:
                    try:wait=(parsedate_to_datetime(value)-datetime.now(timezone.utc)).total_seconds()
                    except (ValueError,TypeError,OverflowError):wait=0
                if wait<float('inf'):delay=max(delay,wait)
                await response.aclose()
            await asyncio.sleep(delay)

    async def aclose(self):
        await self.transport.aclose()


def client_factory(headers=None, timeout=None, auth=None):
    return httpx.AsyncClient(headers=headers, timeout=timeout or 120, auth=auth, follow_redirects=False, trust_env=False,transport=RetryTransport())


def decode_result(result):
    if result.isError:
        detail = ' '.join(c.text for c in result.content if hasattr(c, 'text'))
        raise RuntimeError('MCP tool error: ' + detail)
    if result.structuredContent is not None: return result.structuredContent
    return json.loads(next(c.text for c in result.content if hasattr(c, 'text')))


def verify_bundle(content):
    with zipfile.ZipFile(io.BytesIO(content)) as archive:
        if sum(i.file_size for i in archive.infolist()) > 256 * 1024 * 1024:
            raise ValueError('Expanded artifact exceeds limit')
        metadata = json.loads(archive.read('dataset_metadata.json'))
        hashes = metadata['entries_sha256']
        if not hashes: raise ValueError('Artifact has no checksums')
        for name, expected in hashes.items():
            if hashlib.sha256(archive.read(name)).hexdigest() != expected:
                raise ValueError('Artifact checksum mismatch: ' + name)
        return metadata


async def run(args):
    base = os.environ['ISYNTH_BASE_URL'].rstrip('/')
    parsed = urlsplit(base)
    if parsed.scheme != 'https' and not (parsed.scheme == 'http' and parsed.hostname in {'localhost','127.0.0.1'}):
        raise ValueError('Use HTTPS outside localhost')
    if parsed.username or parsed.password or parsed.query or parsed.fragment or parsed.path:
        raise ValueError('Use a site origin without path, query, fragment or embedded credentials')
    headers = {'Authorization': 'Bearer ' + os.environ['ISYNTH_AUTOMATION_TOKEN']}
    tools = {'identity':'get_identity','schema':'get_schema','search-schema':'get_search_schema','fields':'get_field_catalog','datasets':'list_reaction_datasets','reaction':'get_reaction','statistics':'get_statistics','model':'get_model_status','groups':'get_functional_groups','parameters':'get_parameter_catalog',
        'search':'search_reaction_records','plan':'plan_dataset','build':'build_dataset','status':'get_build','upload':'upload_records','download':'get_build','publication':'get_publication_metadata','start-upload':'start_upload','upload-batch':'upload_batch','upload-status':'get_upload','finish-upload':'finish_upload','cancel-upload':'cancel_upload'}
    tools.update({'parse-report':'get_parse_report','retry-processing':'retry_file_processing','validate':'validate_reaction_records'})
    arguments = {}
    if args.action in {'parse-report','retry-processing'}:arguments['file_id']=str(UUID(args.input or ''))
    if args.action=='parse-report':arguments.update(disposition=args.disposition,page=args.page)
    if args.action in {'search','plan','build','upload','start-upload','upload-batch','validate'}:
        if not args.input: raise ValueError('Supply the input JSON file')
        raw = Path(args.input).read_bytes()
        if len(raw) > 2 * 1024 * 1024: raise ValueError('JSON request exceeds 2 MiB. Use upload_csv.py to stream one complete CSV, or resume batches under one upload_id; never split the dataset.')
        if args.action == 'validate':
            payload = json.loads(raw)
            if isinstance(payload, dict) and 'rows' in payload:
                arguments = {'records': payload['rows'], 'column_mapping': payload.get('column_mapping', {}), 'source_units': payload.get('source_units', {})}
            else:
                arguments = {'records': payload if isinstance(payload, list) else [payload]}
            if not 1 <= len(arguments['records']) <= 50:
                raise ValueError('Validate 1–50 records per call. Split validation calls only; keep the complete source together for upload.')
        elif args.action == 'upload-batch': arguments = json.loads(raw)  # {upload_id, request: {batch_index, rows}}
        else: arguments['query' if args.action == 'search' else 'request'] = json.loads(raw)
    if args.action == 'reaction': arguments['record_id'] = str(UUID(args.input or ''))
    if args.action == 'publication': arguments['doi'] = args.input or ''
    if args.action in {'upload-status','finish-upload','cancel-upload'}: arguments['upload_id'] = str(UUID(args.input or ''))
    if args.action in {'status','download'}:
        arguments['run_id'] = str(UUID(args.input or ''))
    async with streamablehttp_client(base + '/api/mcp', headers=headers, timeout=120,
            httpx_client_factory=client_factory) as (read, write, _session_id):
        async with ClientSession(read, write) as session:
            await session.initialize()
            if args.action == 'tools':
                result = {'tools':[t.model_dump(mode='json') for t in (await session.list_tools()).tools]}
            else: result = decode_result(await session.call_tool(tools[args.action], arguments))
    if args.action == 'download':
        if not args.output: raise ValueError('Provide --output dataset.zip')
        expected_path = '/api/automation/v1/runs/' + arguments['run_id'] + '/bundle'
        if result.get('artifact_path') != expected_path: raise ValueError('Build artifact is not ready or has an unexpected path')
        async with client_factory(headers=headers) as client:
            async with client.stream('GET', base + expected_path) as response:
                response.raise_for_status(); chunks=[];size=0
                async for chunk in response.aiter_bytes():
                    size += len(chunk)
                    if size > 64 * 1024 * 1024: raise ValueError('Artifact exceeds 64 MB')
                    chunks.append(chunk)
        content=b''.join(chunks);metadata=verify_bundle(content)
        with open(args.output,'xb') as target: target.write(content)
        print(json.dumps({'saved':args.output,'checksums_verified':len(metadata['entries_sha256'])}));return
    encoded=json.dumps(result,ensure_ascii=False,indent=2)
    if args.output:
        with open(args.output,'x',encoding='utf-8') as target:target.write(encoded+'\n')
        print('Saved',args.output)
    else: print(encoded)


def main():
    parser=argparse.ArgumentParser(description=__doc__)
    parser.add_argument('action',choices=['identity','tools','schema','search-schema','fields','datasets','reaction','statistics','model','groups','parameters','search','plan','build','status','upload','download','publication','start-upload','upload-batch','upload-status','finish-upload','cancel-upload','parse-report','retry-processing','validate'])
    parser.add_argument('input',nargs='?');parser.add_argument('--output')
    parser.add_argument('--disposition',choices=['rejected','all','accepted','ignored'],default='rejected')
    parser.add_argument('--page',type=int,default=1)
    args=parser.parse_args()
    try: asyncio.run(run(args))
    except (ValueError,KeyError,RuntimeError,httpx.HTTPError,FileExistsError) as exc: raise SystemExit(str(exc)) from None


if __name__ == '__main__': main()
