import copy
import datetime as dt
import json
import itertools
from pathlib import Path
import unittest

from odexa_ref import contracts as c
from profiles import storage_sessions as s

T0 = '2026-09-17T00:00:00Z'
def at(seconds): return (s._time(T0) + dt.timedelta(seconds=seconds)).strftime('%Y-%m-%dT%H:%M:%SZ')
def uid(n): return f'00000000-0000-4000-8000-{n:012d}'
B = dict(protocol_version=s.VERSION, origin='https://publisher.example', service_id='https://publisher.example/services/a',
         issuer='https://provider.example/operator', delegation_id='https://publisher.example/delegations/a')
TERMS = dict(agreement_id=uid(99), report_deadline_seconds=300, retention_seconds=864000, use_expires_at=at(1728000))


def action(n=1, *, action='store', when=0, parent=None):
    cc = dict(copy_id=uid(n+1000), acquired_at=T0, retain_until=at(864000), parent_write=parent) if action == 'store' else None
    return dict(B, event_profile=s.PROFILE, event_id=uid(n), event_type='use.reported', occurred_at=at(when),
        reporter_id='https://agent.example/reporter', source='client_reported', resource_url=B['origin']+'/about',
        policy_id=B['origin']+'/policy', policy_revision=1, actions=[action], purposes=['internal_knowledge'], agreement_id=uid(99),
        asset_ref=None, operation=dict(id=uid(n+100), started_at=at(when), ended_at=at(when), state='completed', kind='action',
        copy_context=cc, storage=None), http=None, quantity=None, unit=None, related_events=[], derived_from=[])


def ref(event, raw=None):
    return dict(reporter_id=event['reporter_id'], event_id=event['event_id'], payload_digest=c.digest(raw or c.json_bytes(event)))


def session(write, n=10, *, trigger='start', index=0, when=None, session_id=uid(500)):
    value = copy.deepcopy(write); start = write['occurred_at']
    if when is None: when = int((s._time(start)-s._time(T0)).total_seconds())
    value.update(event_id=uid(n), occurred_at=at(when), related_events=[ref(write)])
    value['operation'] = dict(id=session_id, started_at=start, ended_at=at(when) if trigger == 'cessation' else None,
        state='completed' if trigger == 'cessation' else 'in_progress', kind='storage_session',
        copy_context=copy.deepcopy(write['operation']['copy_context']), storage=dict(trigger=trigger,
        checkpoint_index=None if trigger == 'cessation' else index,
        scheduled_at=at(when) if trigger != 'checkpoint' else s._stamp(s._time(start)+dt.timedelta(seconds=index*86400))))
    return value


def record(event, received=None, raw=None):
    return dict(payload=raw or c.json_bytes(event), received_at=received or event['occurred_at'], authenticated_reporter_id=event['reporter_id'])


def reduce(events, as_of=at(172801), terms=None):
    return s.reduce_records([v if 'payload' in v else record(v) for v in events], as_of=as_of, terms=terms or TERMS, expected_binding=B)


class StorageSessionTests(unittest.TestCase):
    def test_complete_write_many_reports_one_action(self):
        w=action(); start=session(w); one=session(w,11,trigger='checkpoint',index=1,when=86400)
        end=session(w,12,trigger='cessation',when=100000)
        r=reduce([w,start,one,end]); self.assertEqual(r['completed_action_claims'],1)
        self.assertEqual((r['sessions_started'],r['sessions_closed'],r['checkpoints_expected'],r['checkpoints_timely']),(1,1,1,1))

    def test_additive_discriminator_and_closed_operation(self):
        w=action()
        variants=[]
        v=copy.deepcopy(w);del v['event_profile'];variants.append(v)
        v=copy.deepcopy(w);v['operation']['kind']='other';variants.append(v)
        v=copy.deepcopy(w);v['operation']['extra']=1;variants.append(v)
        v=copy.deepcopy(w);v['operation']['storage']={};variants.append(v)
        v=copy.deepcopy(w);v['event_profile']='storage_sessions_v2';variants.append(v)
        for v in variants:
            with self.subTest(v=v),self.assertRaises(ValueError):s.validate_event(v)

    def test_wrong_session_action_state_or_counter_rejected(self):
        w=action();base=session(w)
        variants=[]
        for field,value in [('actions',['retrieve']),('source','origin_observed')]:
            v=copy.deepcopy(base);v[field]=value;variants.append(v)
        v=copy.deepcopy(base);v['operation']['state']='failed';v['operation']['ended_at']=T0;variants.append(v)
        v=copy.deepcopy(base);v['operation']['storage']['checkpoint_index']=False;variants.append(v)
        v=copy.deepcopy(base);v['quantity']=1;v['unit']='use';variants.append(v)
        for v in variants:
            with self.subTest(v=v),self.assertRaises(ValueError):s.validate_event(v)

    def test_exact_retry_not_new_action_session_or_checkpoint(self):
        w=action();a=session(w);p=session(w,11,trigger='checkpoint',index=1,when=86400)
        r=reduce([w,a,p,record(p,at(87000)),w,a]);self.assertEqual(r['exact_retries'],3)
        self.assertEqual((r['completed_action_claims'],r['sessions_known'],r['checkpoints_timely']),(1,1,1))

    def test_same_event_id_changed_bytes_quarantines_operation(self):
        w=action();a=session(w);changed=copy.deepcopy(a);changed['purposes']=['public_retrieval']
        r=reduce([w,a,changed]);self.assertEqual(r['event_conflicts'],1)
        self.assertEqual(r['unresolved_sessions'],1);self.assertEqual(r['checkpoints_expected'],0)

    def test_same_event_id_write_and_session_variants_preserve_cohort_in_any_order(self):
        w=action();a=session(w);a['event_id']=w['event_id']
        one,two=reduce([w,a]),reduce([a,w])
        self.assertEqual(one,two)
        self.assertEqual((one['event_conflicts'],one['operation_conflicts'],one['completed_action_claims']),(1,2,0))
        self.assertEqual((one['sessions_known'],one['unresolved_sessions']),(1,1))

    def test_all_conflict_variants_and_exact_retries_are_order_independent(self):
        w=action();a=session(w);a['event_id']=w['event_id']
        b=session(w,session_id=uid(501));b['event_id']=w['event_id']
        expected=reduce([w,a,b,a])
        for permutation in itertools.permutations([w,a,b,a]):
            self.assertEqual(reduce(list(permutation)),expected)
        self.assertEqual(expected['exact_retries'],1)
        self.assertEqual((expected['sessions_known'],expected['unresolved_sessions'],expected['operation_conflicts']),(2,2,3))
        self.assertEqual(expected['completed_action_claims'],0)

    def test_conflicting_write_variant_is_never_a_resolvable_session_reference(self):
        w=action();changed=copy.deepcopy(w);changed['operation']['id']=uid(1010)
        a=session(w)
        r=reduce([w,changed,a]);self.assertEqual(r['completed_action_claims'],0)
        self.assertEqual(r['unresolved_sessions'],1)
        self.assertIn('unresolved_or_invalid_write',r['sessions'][0]['reasons'])

    def test_different_event_id_same_checkpoint_slot_conflicts(self):
        w=action();a=session(w);p=session(w,11,trigger='checkpoint',index=1,when=86400)
        q=copy.deepcopy(p);q['event_id']=uid(12)
        r=reduce([w,a,p,q]);self.assertEqual(r['unresolved_sessions'],1)

    def test_context_mutation_terminal_restart_and_second_session_conflict(self):
        w=action();a=session(w);end=session(w,11,trigger='cessation',when=100)
        p=session(w,12,trigger='checkpoint',index=1,when=86400)
        self.assertEqual(reduce([w,a,end,p])['unresolved_sessions'],1)
        changed=copy.deepcopy(a);changed['event_id']=uid(14);changed['resource_url']=B['origin']+'/other'
        self.assertEqual(reduce([w,a,changed])['operation_conflicts'],1)
        other=session(w,15,session_id=uid(501))
        r=reduce([w,a,other]);self.assertEqual(r['operation_conflicts'],2);self.assertEqual(r['unresolved_sessions'],2)

    def test_action_session_id_collision_is_order_independently_unresolved(self):
        w=action();a=session(w,session_id=w['operation']['id'])
        one,two=reduce([w,a]),reduce([a,w])
        self.assertEqual(one,two);self.assertEqual(one['completed_action_claims'],0)
        self.assertEqual(one['sessions_known'],1);self.assertEqual(one['unresolved_sessions'],1)

    def test_exact_write_digest_and_scope_required(self):
        w=action();a=session(w)
        for change in ('missing','digest','scope','copy'):
            q=copy.deepcopy(a)
            if change=='digest':q['related_events'][0]['payload_digest']='sha256:'+'0'*64
            if change=='scope':q['agreement_id']=uid(98)
            if change=='copy':q['operation']['copy_context']['copy_id']=uid(900)
            r=reduce(([w] if change!='missing' else [])+[q])
            self.assertEqual(r['completed_action_claims'],0 if change=='missing' else 1)
            self.assertEqual(r['unresolved_sessions']+r['invalid_records'],1)

    def test_authenticated_reporter_mismatch_and_future_intake_excluded(self):
        w=action();r=record(w);r['authenticated_reporter_id']='https://other.example/agent'
        out=reduce([r,record(w,at(200000))]);self.assertEqual(out['invalid_records'],1);self.assertEqual(out['future_records'],1)
        self.assertEqual(out['completed_action_claims'],0)

    def test_missing_denominator_includes_known_session_without_start_report(self):
        w=action();p=session(w,11,trigger='checkpoint',index=2,when=172800)
        r=reduce([w,p]);self.assertEqual(r['checkpoints_expected'],2)
        self.assertEqual((r['checkpoints_missing'],r['checkpoints_timely']),(1,1));self.assertTrue(r['sessions'][0]['missing_start'])
        self.assertEqual(r['sessions'][0]['missing_index_ranges'],[[1,1]])

    def test_deadline_exact_boundary_late_and_pending(self):
        w=action();a=session(w);p=session(w,11,trigger='checkpoint',index=1,when=86400)
        r=reduce([w,a,record(p,at(86700))]);self.assertEqual(r['checkpoints_timely'],1)
        r=reduce([w,a,record(p,at(86701))]);self.assertEqual(r['checkpoints_late'],1)
        r=reduce([w,a],at(86700));self.assertEqual(r['checkpoints_pending'],1);self.assertEqual(r['checkpoints_missing'],0)
        r=reduce([w,a],at(86701));self.assertEqual(r['checkpoints_missing'],1)

    def test_late_wake_does_not_backfill_previous_indices(self):
        w=action();a=session(w);p=session(w,11,trigger='checkpoint',index=3,when=259601)
        r=reduce([w,a,p],at(259602));self.assertEqual(r['checkpoints_expected'],3)
        self.assertEqual(r['checkpoints_missing'],2);self.assertEqual(r['checkpoints_late'],1)
        q=copy.deepcopy(p);q['operation']['storage'].update(checkpoint_index=1,scheduled_at=at(86400))
        with self.assertRaises(ValueError):s.validate_event(q)

    def test_cessation_boundary_discharges_checkpoint_but_not_earlier_gaps(self):
        w=action();a=session(w);end=session(w,12,trigger='cessation',when=172800)
        r=reduce([w,a,end]);self.assertEqual(r['checkpoints_expected'],1);self.assertEqual(r['checkpoints_missing'],1)
        p=session(w,11,trigger='checkpoint',index=2,when=172800)
        self.assertEqual(reduce([w,a,p,end])['unresolved_sessions'],1)

    def test_failed_cleanup_does_not_fabricate_cessation_or_drop_expected_checks(self):
        w=action();a=session(w)
        r=reduce([w,a],at(864001));self.assertTrue(r['sessions'][0]['retention_overdue'])
        self.assertEqual(r['sessions_closed'],0);self.assertEqual(r['checkpoints_expected'],10)

    def test_copy_inherits_time_and_resolves_parent_exactly(self):
        w=action();a=session(w);child=action(2,when=100,parent=ref(w));b=session(child,20,session_id=uid(600))
        r=reduce([w,a,child,b]);self.assertEqual(r['sessions_started'],2);self.assertEqual(r['completed_action_claims'],2)
        r=reduce([child,b]);self.assertEqual(r['unresolved_sessions'],1)
        reset=copy.deepcopy(child);reset['operation']['copy_context'].update(acquired_at=at(100),retain_until=at(864100))
        bs=session(reset,20,session_id=uid(600));r=reduce([w,reset,bs]);self.assertEqual(r['unresolved_sessions'],1)

    def test_wrong_claimed_retention_deadline_invalid_and_original_use_expiry_caps(self):
        w=action();w['operation']['copy_context']['retain_until']=at(864001)
        self.assertEqual(reduce([w])['invalid_records'],1)
        terms=dict(TERMS,use_expires_at=at(100));w=action();w['operation']['copy_context']['retain_until']=at(100)
        r=reduce([w,session(w)],at(101),terms);self.assertTrue(r['sessions'][0]['retention_overdue'])

    def test_zero_retention_claim_is_visible_overdue_not_authority_to_store(self):
        terms=dict(TERMS,retention_seconds=0);w=action();w['operation']['copy_context']['retain_until']=T0
        r=reduce([w,session(w)],at(1),terms);self.assertTrue(r['sessions'][0]['retention_overdue'])
        self.assertEqual(r['completed_action_claims'],1)

    def test_retrieval_and_failed_write_do_not_start_storage(self):
        w=action(action='retrieve');self.assertEqual(reduce([w])['completed_action_claims'],1)
        w=action();w['operation']['state']='failed';a=session(w)
        r=reduce([w,a]);self.assertEqual(r['completed_action_claims'],0);self.assertEqual(r['unresolved_sessions'],1)

    def test_record_order_is_irrelevant_and_old_payload_whitespace_is_not_equivalent(self):
        w=action();a=session(w);end=session(w,12,trigger='cessation',when=100)
        self.assertEqual(reduce([w,a,end]),reduce([end,a,w]))
        r=reduce([record(w,raw=json.dumps(w,indent=1).encode()),a]);self.assertEqual(r['unresolved_sessions'],1)

    def test_large_time_span_has_compressed_missing_ranges(self):
        terms=dict(TERMS,retention_seconds=None,use_expires_at=None);w=action();w['operation']['copy_context']['retain_until']=None
        r=reduce([w,session(w)],'9999-01-01T00:00:00Z',terms)
        self.assertGreater(r['checkpoints_expected'],1000000)
        self.assertLess(len(json.dumps(r)),4000)

    def test_large_retention_still_uses_earlier_accepted_expiry(self):
        terms=dict(TERMS,retention_seconds=9007199254740991,use_expires_at=at(100))
        w=action();w['operation']['copy_context']['retain_until']=at(100)
        r=reduce([w,session(w)],at(101),terms);self.assertEqual(r['invalid_records'],0)
        self.assertTrue(r['sessions'][0]['retention_overdue'])

    def test_large_reporting_deadline_does_not_overflow_or_invent_missing(self):
        terms=dict(TERMS,report_deadline_seconds=9007199254740991)
        w=action();a=session(w);p=session(w,11,trigger='checkpoint',index=1,when=86400)
        r=reduce([w,a,p],at(172801),terms);self.assertEqual(r['checkpoints_timely'],1)
        self.assertEqual(r['checkpoints_pending'],1);self.assertEqual(r['checkpoints_missing'],0)
        self.assertIsNone(r['sessions'][0]['checkpoints'][0]['deadline_at'])
        self.assertTrue(r['sessions'][0]['checkpoints'][0]['deadline_after_supported_calendar'])


def cases():
    w=action();a=session(w);p=session(w,11,trigger='checkpoint',index=1,when=86400);end=session(w,12,trigger='cessation',when=100000)
    good=[w,a,p,end,action(2,action='retrieve')]
    bad=[]
    for label, mutate in [
        ('legacy_no_discriminator',lambda v:v.pop('event_profile')),
        ('unknown_kind',lambda v:v['operation'].update(kind='unknown')),
        ('false_index',lambda v:v['operation']['storage'].update(checkpoint_index=False)),
        ('session_failed',lambda v:v['operation'].update(state='failed',ended_at=T0)),
        ('extra_operation_property',lambda v:v['operation'].update(extra=1)),
        ('session_quantity',lambda v:v.update(quantity=1,unit='billable')),
    ]:
        v=copy.deepcopy(a);mutate(v);bad.append(dict(name=label,event=v,valid=False))
    return [dict(name='positive_'+str(i),event=v,valid=True) for i,v in enumerate(good)]+bad


if __name__=='__main__':
    import sys
    if '--cases' in sys.argv: print(json.dumps(cases(),indent=2))
    else: unittest.main()
