# shellcheck shell=dash

___x_cmd_claude_session_extract(){
    local session_id=""
    local output_dir=""
    while [ "$#" -gt 0 ]; do
        case "$1" in
            --help|-h)  ___x_cmd help -m claude session extract "$@" ; return 0 ;;
            --session)  session_id="$2"; arg:2:shift ;;
            --output)   output_dir="$2"; arg:2:shift ;;
            *) ;;
        esac
    done

    [ -n "$session_id" ] || N=claude M="Please provide session id" log:ret:64
    [ -n "$output_dir" ] || output_dir="$___X_CMD_CLAUDE_ROOT_PATH/session/extract/$session_id"

    local session_file; session_file="$(___x_cmd_cmds find "$___X_CMD_CLAUDE_CODE_DATA_PATH" -name "${session_id}.jsonl" -print -quit)"
    [ -f "$session_file" ] || N=claude M="Cannot find the history conversation file -> $session_id" log:ret:1

    ___x_cmd mkdirp "$output_dir" || N=claude M="Failed to create output directory -> $output_dir" log:ret:1
    ___x_cmd_claude_session_extract___run "$session_file" "$output_dir" "$session_id" || {
        claude:error "Failed to extract session -> $session_id"
        return 1
    }

    claude:info --session_id "$session_id" --output_dir "$output_dir" "Session successfully extracted"
}

___x_cmd_claude_session_extract___run(){
    local session_file="$1"
    local output_dir="$2"
    local session_id="$3"

    ___x_cmd python - "$session_file" "$output_dir" "$session_id" <<'PYEOF' || return $?

import json
import sys
import os
import re
from datetime import datetime, timezone


session_file = sys.argv[1]
output_dir = sys.argv[2]
session_id = sys.argv[3] if len(sys.argv) > 3 else ""

# Lookup pids from sessions directory
session_pids = []
sessions_dir = os.path.expanduser("~/.claude/sessions")
if os.path.isdir(sessions_dir):
    for filename in os.listdir(sessions_dir):
        if not filename.endswith(".json"):
            continue
        filepath = os.path.join(sessions_dir, filename)
        try:
            with open(filepath, "r", encoding="utf-8") as f:
                data = json.load(f)
            if data.get("sessionId") == session_id:
                pid_num = os.path.splitext(filename)[0]
                session_pids.append(pid_num)
        except (json.JSONDecodeError, OSError):
            continue


def format_timestamp(ts):
    """Convert ISO 8601 timestamp to Unix timestamp."""
    if not ts:
        return ""
    ts = ts.replace('Z', '+00:00')
    ts = re.sub(r'\.\d+', '', ts)
    try:
        dt = datetime.fromisoformat(ts)
        if dt.tzinfo is None:
            dt = dt.replace(tzinfo=timezone.utc)
        return str(int(dt.timestamp()))
    except (ValueError, OSError):
        return ""


TOOLS_FIELDS = ['start_time', 'model', 'cwd', 'tool_name', 'tool_use_id', 'tool_input', 'tool_result', 'is_error', 'stdout', 'stderr', 'interrupted']
TOKENS_FIELDS = ['start_time', 'model', 'input_tokens', 'output_tokens', 'cache_creation_input_tokens', 'cache_read_input_tokens', 'reasoning_output_tokens']
TOKENS_TOTAL_FIELDS = ['total_input_tokens', 'total_output_tokens', 'total_cache_creation_input_tokens', 'total_cache_read_input_tokens', 'total_tokens', 'assistant_message_count']
TURNS_FIELDS = ['turn', 'start_time', 'dir', 'agent', 'user_messages', 'assistant_messages', 'tool_use_count', 'input_tokens', 'output_tokens', 'start_time', 'end_time', 'stop_reason']

def escape_tsv_field(v):
    if v is None:
        return ''
    s = str(v)
    s = s.replace('\\', '\\\\')
    s = s.replace('\t', '\\t')
    s = s.replace('\n', '\\n')
    s = s.replace('\r', '\\r')
    return s

def prep_tool_records(records):
    out = []
    for r in records:
        row = dict(r)
        row['tool_input'] = json.dumps(row.get('tool_input', ''), ensure_ascii=False) if row.get('tool_input', '') else ''
        out.append(row)
    return out

def _write(directory, filename, content):
    with open(os.path.join(directory, filename), 'w', encoding='utf-8') as f:
        f.write(content)

def write_tsv(records, filepath, fields):
    lines = []
    for r in records:
        row = []
        for f in fields:
            v = r.get(f, '')
            row.append(escape_tsv_field(v))
        lines.append('\t'.join(row))
    with open(filepath, 'w', encoding='utf-8') as f:
        f.write('\n'.join(lines) + '\n')

def write_schema(schema_dir, name, fields):
    os.makedirs(schema_dir, exist_ok=True)
    with open(os.path.join(schema_dir, f'{name}.schema.tsv'), 'w', encoding='utf-8') as f:
        f.write('\t'.join(fields) + '\n')


def is_tool_result_user(record):
    """Check if a user record is a tool result (not initial user input)."""
    if record.get('type') != 'user':
        return False
    msg = record.get('message', {})
    content = msg.get('content', [])
    if isinstance(content, list) and len(content) > 0:
        return all(item.get('type') == 'tool_result' for item in content if isinstance(item, dict))
    return False


def is_local_command_user(record):
    """Check if a user record is a local command message (not real user input)."""
    if record.get('type') != 'user':
        return False
    msg = record.get('message', {})
    content = msg.get('content', '')
    if isinstance(content, str) and content.strip().startswith(('<local-command', '<command-name')):
        return True
    return False

def is_initial_user(record):
    """Check if a user record is the initial user input (start of a turn)."""
    if record.get('type') != 'user':
        return False
    return not is_tool_result_user(record) and not is_local_command_user(record)


with open(session_file, 'r', encoding='utf-8') as f:
    records = [json.loads(line) for line in f if line.strip()]

# Build turn index: assign each record to a turn number
# A turn starts with the first user message for a given promptId.
# Skill injections re-use the original promptId and must not create new turns.
turn_records = []  # list of (turn_num, record)
current_turn = 0
seen_prompt_ids = set()

for r in records:
    if is_initial_user(r):
        pid = r.get('promptId', '')
        if pid not in seen_prompt_ids:
            seen_prompt_ids.add(pid)
            current_turn += 1
    turn_records.append((current_turn, r))

# Extract metadata from first meaningful record
meta = {}
first_assistant_model = ''
for r in records:
    if r.get('type') in ('user', 'assistant'):
        meta = {
            'session_id': session_id,
            'pids': session_pids,
            'cwd': r.get('cwd', ''),
            'git_branch': r.get('gitBranch', ''),
            'start_time': format_timestamp(r.get('timestamp', '')),
            'end_time': format_timestamp(next((r.get('timestamp', '') for r in reversed(records) if r.get('timestamp')), '') if records else ''),
            'record_count': len(records),
            'turn_count': current_turn,
            'user_count': sum(1 for r in records if r.get('type') == 'user'),
            'assistant_count': sum(1 for r in records if r.get('type') == 'assistant'),
            'tool_use_count': sum(
                1 for r in records if r.get('type') == 'assistant'
                for item in r.get('message', {}).get('content', [])
                if isinstance(item, dict) and item.get('type') == 'tool_use'
            )
        }
        break

# Extract model from first assistant record
for r in records:
    if r.get('type') == 'assistant':
        first_assistant_model = r.get('message', {}).get('model', '')
        break

if records:
    is_complete = records[-1].get('type') == 'last-prompt'
    meta['is_complete'] = str(is_complete).lower()
else:
    meta['is_complete'] = ''

meta['model'] = first_assistant_model
meta['platform'] = 'claude'

meta_dir = os.path.join(output_dir, 'metadata')
os.makedirs(meta_dir, exist_ok=True)
for key, val in meta.items():
    with open(os.path.join(meta_dir, key), 'w', encoding='utf-8') as f:
        if isinstance(val, list):
            f.write('\n'.join(str(v) for v in val) + '\n')
        else:
            f.write(str(val))

# Index tool_results by tool_use_id for pairing
tool_results = {}
for r in records:
    if r.get('type') != 'user':
        continue
    msg = r.get('message', {})
    content = msg.get('content', [])
    if isinstance(content, list):
        for item in content:
            if isinstance(item, dict) and item.get('type') == 'tool_result':
                tuid = item.get('tool_use_id', '')
                tur = r.get('toolUseResult', {})
                if isinstance(tur, str):
                    tur = {'stdout': '', 'stderr': tur, 'interrupted': False, 'isImage': False, 'noOutputExpected': False}
                tool_results[tuid] = {
                    'content': item.get('content', ''),
                    'is_error': item.get('is_error', False),
                    'toolUseResult': tur,
                    'end_time': format_timestamp(r.get('timestamp', ''))
                }

# Process each turn
turns_summary = []

for turn_num in range(1, current_turn + 1):
    turn_recs = [r for tnum, r in turn_records if tnum == turn_num]
    if not turn_recs:
        continue

    # Find initial user timestamp for directory name
    initial_ts = ''
    for r in turn_recs:
        if is_initial_user(r):
            initial_ts = format_timestamp(r.get('timestamp', ''))
            break

    # Find last timestamp in this turn for end_time
    turn_end_ts = ''
    for r in reversed(turn_recs):
        ets = format_timestamp(r.get('timestamp', ''))
        if ets:
            turn_end_ts = ets
            break

    turn_dir_name = initial_ts
    turn_dir = os.path.join(output_dir, 'turn', turn_dir_name)
    if os.path.isdir(turn_dir):
        continue
    os.makedirs(turn_dir, exist_ok=True)

    # Build turn content
    turn_tools = []
    turn_msgs = []  # list of dicts: {seq, type, timestamp, content, ...}
    turn_tokens = []

    for r in turn_recs:
        t = r.get('type')
        ts = r.get('timestamp', '')
        fts = format_timestamp(ts)

        if t == 'user':
            msg = r.get('message', {})
            content = msg.get('content', '')
            if isinstance(content, str) and content.strip():
                turn_msgs.append({'type': 'user', 'timestamp': fts, 'content': content.strip()})
            elif isinstance(content, list):
                texts = [item.get('text', '') for item in content if isinstance(item, dict) and item.get('type') == 'text']
                text = '\n'.join(t for t in texts if t.strip())
                if text.strip():
                    turn_msgs.append({'type': 'user', 'timestamp': fts, 'content': text.strip()})

        elif t == 'attachment':
            att = r.get('attachment', {})
            att_type = att.get('type', '')
            if att_type.startswith('hook_'):
                hook = {
                    'type': 'hook', 'timestamp': fts,
                    'hook_event': att.get('hookEvent', ''),
                    'hook_name': att.get('hookName', ''),
                    'hook_command': att.get('command', ''),
                }
                if att.get('toolUseID'):
                    hook['tool_use_id'] = att['toolUseID']
                if att.get('exitCode') is not None:
                    hook['hook_exit_code'] = str(att['exitCode'])
                if att.get('durationMs') is not None:
                    hook['hook_duration_ms'] = str(att['durationMs'])
                if att.get('stdout'):
                    hook['hook_stdout'] = att['stdout']
                if att.get('stderr'):
                    hook['hook_stderr'] = att['stderr']
                turn_msgs.append(hook)

        elif t == 'assistant':
            msg = r.get('message', {})
            model = msg.get('model', '')
            content = msg.get('content', [])
            turn_assistant_texts = []

            for item in content:
                if not isinstance(item, dict):
                    continue
                itype = item.get('type')
                if itype == 'text':
                    txt = item.get('text', '')
                    if txt.strip():
                        turn_assistant_texts.append(txt)
                elif itype == 'thinking':
                    th = item.get('thinking', '')
                    if th.strip():
                        turn_msgs.append({'type': 'thinking', 'timestamp': fts, 'content': th.strip()})
                elif itype == 'tool_use':
                    tuid = item.get('id', '')
                    tname = item.get('name', 'unknown')
                    tinput = item.get('input', {})
                    tres = tool_results.get(tuid, {})
                    tr = {
                        'start_time': format_timestamp(ts),
                        'model': model,
                        'cwd': r.get('cwd', ''),
                        'tool_name': tname,
                        'tool_use_id': tuid,
                        'tool_input': tinput,
                        'tool_result': tres.get('content', ''),
                        'is_error': tres.get('is_error', False),
                        'stdout': tres.get('toolUseResult', {}).get('stdout', ''),
                        'stderr': tres.get('toolUseResult', {}).get('stderr', ''),
                        'interrupted': str(tres.get('toolUseResult', {}).get('interrupted')).lower() if 'interrupted' in (tres.get('toolUseResult') or {}) else ''
                    }
                    turn_tools.append(tr)

                    tres_content = tres.get('content', '')
                    if isinstance(tres_content, list):
                        tres_content = '\n'.join(
                            item.get('text', '') for item in tres_content
                            if isinstance(item, dict) and item.get('type') == 'text'
                        )
                    turn_msgs.append({
                        'type': 'tool_use', 'timestamp': fts,
                        'tool_name': tname, 'tool_use_id': tuid,
                        'tool_input': tinput, 'model': model,
                        'cwd': r.get('cwd', ''),
                        'tool_result_text': tres_content if isinstance(tres_content, str) else '',
                        'is_error': tres.get('is_error', False),
                        'stdout': tres.get('toolUseResult', {}).get('stdout', ''),
                        'stderr': tres.get('toolUseResult', {}).get('stderr', ''),
                        'end_time': tres.get('end_time', ''),
                        'interrupted': str(tres.get('toolUseResult', {}).get('interrupted')) if 'interrupted' in (tres.get('toolUseResult') or {}) else '',
                    })

            if turn_assistant_texts:
                turn_msgs.append({'type': 'assistant', 'timestamp': fts, 'model': model, 'content': '\n\n'.join(turn_assistant_texts)})

            # Token usage
            usage = msg.get('usage', {})
            if usage:
                turn_tokens.append({
                    'start_time': format_timestamp(ts),
                    'model': model,
                    'input_tokens': usage.get('input_tokens', 0),
                    'output_tokens': usage.get('output_tokens', 0),
                    'cache_creation_input_tokens': usage.get('cache_creation_input_tokens', ''),
                    'cache_read_input_tokens': usage.get('cache_read_input_tokens', ''),
                    'reasoning_output_tokens': ''
                })


    # Write turn files
    turn_cwd = ''
    turn_git_branch = ''
    for r in turn_recs:
        if r.get('type') in ('user', 'assistant'):
            turn_cwd = r.get('cwd', '')
            turn_git_branch = r.get('gitBranch', '')
            break

    _write(turn_dir, 'agent', 'claude')
    _write(turn_dir, 'start_time', initial_ts)
    _write(turn_dir, 'end_time', turn_end_ts)
    _write(turn_dir, 'cwd', turn_cwd)
    _write(turn_dir, 'git_branch', turn_git_branch)

    write_tsv(prep_tool_records(turn_tools), os.path.join(turn_dir, 'tools.tsv'), TOOLS_FIELDS)

    # Write individual message directories
    msg_dir = os.path.join(turn_dir, 'msg')
    os.makedirs(msg_dir, exist_ok=True)
    msg_width = len(str(len(turn_msgs))) if turn_msgs else 1
    for idx, m in enumerate(turn_msgs, 1):
        mtype = m['type']
        seq = str(idx).zfill(msg_width)

        if mtype == 'tool_use':
            # NNN_tool/ directory
            tdir = os.path.join(msg_dir, f'{seq}_tool')
            os.makedirs(tdir, exist_ok=True)
            _write(tdir, 'name', m['tool_name'])
            _write(tdir, 'tool_use_id', m['tool_use_id'])
            with open(os.path.join(tdir, 'input.json'), 'w', encoding='utf-8') as f:
                json.dump(m['tool_input'], f, ensure_ascii=False, indent=2)
            _write(tdir, 'start_time', m.get('timestamp', ''))
            _write(tdir, 'model', m.get('model', ''))
            _write(tdir, 'cwd', m.get('cwd', ''))
            if m['tool_name'] == 'Agent' and isinstance(m.get('tool_input'), dict):
                _write(tdir, 'prompt.md', m['tool_input'].get('prompt', ''))
                _write(tdir, 'subagent_type', m['tool_input'].get('subagent_type', ''))
                _write(tdir, 'description', m['tool_input'].get('description', ''))

            # NNN_tool_result/ directory (separate from tool)
            tr_dir = os.path.join(msg_dir, f'{seq}_tool_result')
            os.makedirs(tr_dir, exist_ok=True)
            _write(tr_dir, 'tool_use_id', m['tool_use_id'])
            _write(tr_dir, 'result.md', m.get('tool_result_text', ''))
            if m.get('end_time'):
                _write(tr_dir, 'end_time', m['end_time'])
            if m.get('stdout'):
                _write(tr_dir, 'stdout.md', m['stdout'])
            if m.get('stderr'):
                _write(tr_dir, 'stderr.md', m['stderr'])
            if m.get('is_error'):
                _write(tr_dir, 'error', str(m['is_error']).lower())
            if m.get('interrupted'):
                _write(tr_dir, 'interrupted', m['interrupted'])

        elif mtype == 'hook':
            # NNN_hook/ directory
            hdir = os.path.join(msg_dir, f'{seq}_hook')
            os.makedirs(hdir, exist_ok=True)
            _write(hdir, 'hook_event', m['hook_event'])
            _write(hdir, 'hook_name', m['hook_name'])
            _write(hdir, 'hook_command', m['hook_command'])
            if m.get('tool_use_id'):
                _write(hdir, 'tool_use_id', m['tool_use_id'])
            if m.get('hook_exit_code'):
                _write(hdir, 'hook_exit_code', m['hook_exit_code'])
            if m.get('hook_duration_ms'):
                _write(hdir, 'hook_duration_ms', m['hook_duration_ms'])
            if m.get('hook_stdout'):
                _write(hdir, 'hook_stdout.md', m['hook_stdout'])
            if m.get('hook_stderr'):
                _write(hdir, 'hook_stderr.md', m['hook_stderr'])

        else:
            # user/thinking/assistant: directory with content.md + start_time (+ model for assistant)
            mdir = os.path.join(msg_dir, f'{seq}_{mtype}')
            os.makedirs(mdir, exist_ok=True)
            _write(mdir, 'content.md', m['content'])
            _write(mdir, 'start_time', m.get('timestamp', ''))
            if mtype == 'assistant':
                _write(mdir, 'model', m.get('model', ''))

    write_tsv(turn_tokens, os.path.join(turn_dir, 'tokens.tsv'), TOKENS_FIELDS)

    def safe_sum(records, key):
        vals = [r.get(key, 0) for r in records if isinstance(r.get(key), (int, float))]
        return sum(vals) if vals else ''

    tokens_total = {
        'total_input_tokens': safe_sum(turn_tokens, 'input_tokens'),
        'total_output_tokens': safe_sum(turn_tokens, 'output_tokens'),
        'total_cache_creation_input_tokens': safe_sum(turn_tokens, 'cache_creation_input_tokens'),
        'total_cache_read_input_tokens': safe_sum(turn_tokens, 'cache_read_input_tokens'),
        'total_tokens': safe_sum(turn_tokens, 'input_tokens') + safe_sum(turn_tokens, 'output_tokens') if turn_tokens else '',
        'assistant_message_count': len(turn_tokens)
    }
    write_tsv([tokens_total], os.path.join(turn_dir, 'tokens_total.tsv'), TOKENS_TOTAL_FIELDS)

    turn_stop_reason = ''
    for r in reversed(turn_recs):
        if r.get('type') == 'assistant':
            sr = r.get('message', {}).get('stop_reason')
            if sr:
                turn_stop_reason = sr
                break

    turns_summary.append({
        'turn': turn_num,
        'start_time': initial_ts,
        'dir': f"turn/{turn_dir_name}",
        'agent': 'claude',
        'user_messages': len([r for r in turn_recs if r.get('type') == 'user']),
        'assistant_messages': len([r for r in turn_recs if r.get('type') == 'assistant']),
        'tool_use_count': len(turn_tools),
        'input_tokens': sum(t['input_tokens'] for t in turn_tokens) if turn_tokens else '',
        'output_tokens': sum(t['output_tokens'] for t in turn_tokens) if turn_tokens else '',
        'start_time': initial_ts,
        'end_time': turn_end_ts,
        'stop_reason': turn_stop_reason
    })

# Write turns summary
write_tsv(turns_summary, os.path.join(output_dir, 'turns.tsv'), TURNS_FIELDS)

# Write unified schema files
schema_dir = os.path.join(output_dir, 'schema')
write_schema(schema_dir, 'tools', TOOLS_FIELDS)
write_schema(schema_dir, 'tokens', TOKENS_FIELDS)
write_schema(schema_dir, 'tokens_total', TOKENS_TOTAL_FIELDS)
write_schema(schema_dir, 'turns', TURNS_FIELDS)
PYEOF
}
