import copy
import gzip
import json
import sys
import unittest

from cryptography.hazmat.primitives.asymmetric import ec
from odexa_ref import crypto
from profiles import asset_evidence as a
from profiles.authority_transport import AuthoritySnapshot
from profiles.gateway_reports import delivery_report
from test_delegation_profile import fixture,observation,NOW,ORIGIN
from test_free_contracts import fixtures as free_fixtures,uid,later


def asset_fixture(external=True):
    meta,pol,op=fixture(external);service=meta['services'][0]
    service['capabilities'].append('publish_assets');service['signing_keys'][0]['uses'].append(a.TYPE)
    service['scope'].update(asset_id_prefixes=[ORIGIN+'/assets/'],version_id_prefixes=[ORIGIN+'/versions/'])
    for grant in meta['delegations']:
        grant['capabilities'].append('publish_assets')
        grant.update(asset_id_prefixes=[ORIGIN+'/assets/'],version_id_prefixes=[ORIGIN+'/versions/'])
    return meta,pol,op


def snapshot(meta,pol,checked=NOW):
    obs=observation(meta,pol);obs.update(origin=meta['origin'],authority_url=meta['origin']+'/odexa-service.json',
        policy_url=meta['origin']+'/odexa.json',checked_at=checked)
    return AuthoritySnapshot(a.raw(meta),a.raw(pol),obs)


def manifest(meta,pol,op,body=b'originator asset\n',number=1,parents=()):
    service=meta['services'][0];version=meta['origin']+'/versions/'+str(number)
    operation=dict(op,operation_id=uid(100+number),kind='publish_manifest',asset_id=meta['origin']+'/assets/article',
        version_id=version,endpoint=version,key_use=a.TYPE,offer_seconds=0,access_seconds=0,use_seconds=0)
    value=dict(protocol_version=a.VERSION,profile=a.PROFILE,origin=meta['origin'],asset_id=operation['asset_id'],version_id=version,
        published_at=NOW,authority=dict(origin=meta['origin'],service_id=service['id'],revision=meta['revision'],
            url=meta['origin']+'/odexa-service/history/'+str(meta['revision'])+'.json',digest=a.digest(a.raw(meta))),
        representations=[dict(representation_id=uid(number),media_type='text/plain',media_parameters={'charset':'utf-8'},
            languages=['en'],decoded_length=len(body),decoded_digest=a.digest(body))],
        issuer=service['issuer'],delegation_id=operation['delegation_id'],policy_digest=a.digest(a.raw(pol)),publication=operation,
        content_variants=[],derived_from=[dict(representation_id=uid(number),source=p,relationship='copy') for p in parents])
    return value


def ref(value):
    return dict(asset_id=value['asset_id'],version_id=value['version_id'],manifest_digest=a.digest(a.raw(value)),representation_id=value['representations'][0]['representation_id'])


def bundle(values,meta,pol,roots=None,key=None):
    key=key or ec.derive_private_key(1,ec.SECP256R1())
    return a.raw(dict(protocol_version=a.VERSION,profile=a.PROFILE,roots=roots or [ref(values[-1])],
        manifests=[a.sign_manifest(a.raw(v),key,v['publication']['key_id']) for v in values],
        documents=[dict(digest=a.digest(body),body_b64url=crypto.b64u(body)) for body in (a.raw(meta),a.raw(pol))]))


class AssetEvidenceTests(unittest.TestCase):
    def setUp(self):
        self.meta,self.pol,self.op=asset_fixture();self.body=b'originator asset\n'
        self.value=manifest(self.meta,self.pol,self.op,self.body)
        self.snap=snapshot(self.meta,self.pol)
        self.trust=a.AssetTrust({ORIGIN:self.snap},{a.digest(x):x for x in (self.snap.authority_bytes,self.snap.policy_bytes)})
        self.clock=lambda:NOW

    def verify(self,values=None,raw=None,trust=None):
        return a.verify_closure(raw or bundle(values or [self.value],self.meta,self.pol),trust=trust or self.trust,now=self.clock)

    def delivery(self,**changes):
        f,_=free_fixtures();r=self.value['representations'][0]
        options=dict(event_id=uid(500),hop_id=uid(501),started_at=NOW,occurred_at=NOW,
            bytes_written=len(self.body),representation_digest=a.digest(self.body),complete=True,
            representation_metadata={k:r[k] for k in ('media_type','media_parameters','languages')},asset_ref=ref(self.value))
        options.update(changes)
        return delivery_report(f['binding'],f['introspection_response'],ORIGIN+'/gateway',**options)

    def test_signed_direct_and_delegated_manifest_is_evidence_not_permission(self):
        result=self.verify();self.assertFalse(result.resource_permission_granted)
        self.assertTrue(next(iter(result.manifests.values()))['publication_authorized_now'])
        meta,pol,op=asset_fixture(False);m=manifest(meta,pol,op);snap=snapshot(meta,pol)
        direct=a.verify_closure(bundle([m],meta,pol),trust=a.AssetTrust({ORIGIN:snap},{}),now=self.clock)
        self.assertFalse(direct.resource_permission_granted)

    def test_closed_version_signature_and_byte_substitution(self):
        raw=bundle([self.value],self.meta,self.pol);v=a.load(raw,a.MAX_BUNDLE)
        h,p,s=v['manifests'][0].split('.');payload=a.load(crypto.unb64u(p));payload['representations'][0]['decoded_length']+=1
        v['manifests'][0]=h+'.'+crypto.b64u(a.raw(payload))+'.'+s
        with self.assertRaises(a.AssetError):self.verify(raw=a.raw(v))
        for field,value in [('protocol_version','1.2.0-draft.2'),('profile','unknown'),('extra',True)]:
            changed=copy.deepcopy(self.value);changed[field]=value
            with self.assertRaises(ValueError):a.validate_manifest(changed)

    def test_untrusted_bundle_cannot_bootstrap_origin_history(self):
        next_meta=copy.deepcopy(self.meta);next_meta['revision']=2
        current=snapshot(next_meta,self.pol)
        with self.assertRaisesRegex(a.AssetError,'independent exact origin pin'):
            self.verify(trust=a.AssetTrust({ORIGIN:current},{}))
        with self.assertRaises(a.AssetError):self.verify(trust=a.AssetTrust({},self.trust.pinned_documents))

    def test_missing_extra_duplicate_context_and_unreachable_manifest_rejected(self):
        for change in (lambda b:b['documents'].pop(),lambda b:b['documents'].append(b['documents'][0]),
                       lambda b:b['manifests'].append(b['manifests'][0])):
            v=a.load(bundle([self.value],self.meta,self.pol),a.MAX_BUNDLE);change(v)
            with self.assertRaises(a.AssetError):self.verify(raw=a.raw(v))
        other=manifest(self.meta,self.pol,self.op,number=2)
        with self.assertRaisesRegex(a.AssetError,'unreachable'):
            self.verify(raw=bundle([self.value,other],self.meta,self.pol,roots=[ref(self.value)]))

    def test_origin_scope_issuer_key_use_and_endpoint_substitution(self):
        for field,value in [('issuer','https://provider.example/attacker'),('endpoint',ORIGIN+'/versions/other'),
                            ('asset_id',ORIGIN+'/outside/asset'),('key_id','https://provider.example/unknown-key')]:
            m=copy.deepcopy(self.value);m['publication'][field]=value
            if field in ('issuer','asset_id'):m[field]=value
            with self.assertRaises(ValueError):self.verify(values=[m])
        m=copy.deepcopy(self.value);m['publication']['key_use']='odexa-receipt+jws'
        with self.assertRaises(ValueError):a.validate_manifest(m)

    def test_future_or_changed_version_pin_is_rejected(self):
        m=copy.deepcopy(self.value);m['published_at']=later(1)
        with self.assertRaises(a.AssetError):self.verify(values=[m])
        pins={self.value['version_id']:'sha256:'+'0'*64}
        with self.assertRaises(a.AssetError):self.verify(trust=a.AssetTrust(self.trust.current_snapshots,self.trust.pinned_documents,pins))

    def test_retired_signer_preserves_history_but_revoked_material_does_not(self):
        meta=copy.deepcopy(self.meta);meta['revision']=2;key=meta['services'][0]['signing_keys'][0]
        key.update(state='retired',retired_at=NOW)
        trust=a.AssetTrust({ORIGIN:snapshot(meta,self.pol)},self.trust.pinned_documents)
        result=self.verify(trust=trust);entry=next(iter(result.manifests.values()))
        self.assertFalse(entry['publication_authorized_now']);self.assertEqual(entry['signature_assurance'],'historical_origin_authorized')
        key.update(state='revoked',revoked_at=NOW)
        with self.assertRaisesRegex(a.AssetError,'revoked material'):
            self.verify(trust=a.AssetTrust({ORIGIN:snapshot(meta,self.pol)},self.trust.pinned_documents))

    def test_removed_compromised_key_remains_rejected_by_pinned_history(self):
        revoked=copy.deepcopy(self.meta);revoked['revision']=2
        revoked['services'][0]['signing_keys'][0].update(state='revoked',revoked_at=NOW)
        current=copy.deepcopy(self.meta);current['revision']=3
        current['services'][0]['signing_keys'][0].update(kid='https://provider.example/keys/two',public_jwk=crypto.public_jwk(ec.derive_private_key(2,ec.SECP256R1())))
        pins=dict(self.trust.pinned_documents);pins[a.digest(a.raw(revoked))]=a.raw(revoked)
        with self.assertRaisesRegex(a.AssetError,'revoked material'):
            self.verify(trust=a.AssetTrust({ORIGIN:snapshot(current,self.pol)},pins))

    def test_same_key_identity_cannot_be_rebound(self):
        meta=copy.deepcopy(self.meta);meta['revision']=2
        meta['services'][0]['signing_keys'][0]['public_jwk']=crypto.public_jwk(ec.derive_private_key(2,ec.SECP256R1()))
        with self.assertRaisesRegex(a.AssetError,'rebound'):
            self.verify(trust=a.AssetTrust({ORIGIN:snapshot(meta,self.pol)},self.trust.pinned_documents))

    def test_context_expiring_during_verification_and_history_rollback_rejected(self):
        moments=iter([NOW,later(6)])
        with self.assertRaisesRegex(a.AssetError,'stale'):
            a.verify_closure(bundle([self.value],self.meta,self.pol),trust=self.trust,now=lambda:next(moments))
        class Deny:
            def require_current(self,_):raise ValueError('superseded history')
        with self.assertRaises(a.AssetError):
            a.verify_closure(bundle([self.value],self.meta,self.pol),trust=self.trust,now=self.clock,histories={ORIGIN:Deny()})

    def test_exact_copy_derivation_closes_and_changed_copy_is_rejected(self):
        child=manifest(self.meta,self.pol,self.op,self.body,number=2,parents=[ref(self.value)])
        self.assertEqual(len(self.verify(values=[self.value,child]).manifests),2)
        with self.assertRaisesRegex(a.AssetError,'missing referenced'):
            self.verify(values=[child])
        child['representations'][0]['decoded_digest']=a.digest(b'changed')
        with self.assertRaisesRegex(a.AssetError,'copy changed'):
            self.verify(values=[self.value,child])
        child['derived_from'][0]['relationship']='transform'
        self.assertEqual(len(self.verify(values=[self.value,child]).manifests),2)

    def test_depth_bound_and_same_version_conflicts_are_closed(self):
        values=[self.value]
        for n in range(2,10):values.append(manifest(self.meta,self.pol,self.op,self.body,number=n,parents=[ref(values[-1])]))
        self.assertEqual(len(self.verify(values=values[:8]).manifests),8)
        with self.assertRaisesRegex(a.AssetError,'depth exceeded'):self.verify(values=values)
        variant=copy.deepcopy(self.value);variant['representations'][0]['languages']=[]
        with self.assertRaisesRegex(a.AssetError,'version redefined'):self.verify(values=[self.value,variant])

    def test_duplicate_root_or_edge_cannot_hide_behind_object_member_order(self):
        v=a.load(bundle([self.value],self.meta,self.pol),a.MAX_BUNDLE)
        v['roots'].append(dict(reversed(list(v['roots'][0].items()))))
        with self.assertRaisesRegex(a.AssetError,'duplicate root'):self.verify(raw=a.raw(v))
        child=manifest(self.meta,self.pol,self.op,self.body,number=2,parents=[ref(self.value)])
        edge=copy.deepcopy(child['derived_from'][0]);edge['source']=dict(reversed(list(edge['source'].items())))
        child['derived_from'].append(edge)
        with self.assertRaisesRegex(a.AssetError,'duplicate/conflicting'):a.validate_manifest(child)

    def test_signed_digest_claim_is_not_proof_of_client_receipt_or_model_use(self):
        event=self.delivery();closure=self.verify()
        claim=a.correlate_delivery(event,closure)
        self.assertEqual(claim['binding'],'full_representation_consistent');self.assertFalse(claim['representation_bytes_match'])
        matched=a.correlate_delivery(event,closure,content_bytes=self.body)
        self.assertTrue(matched['representation_bytes_match'])
        self.assertFalse(matched['full_version_delivered']);self.assertFalse(matched['downstream_use_verified'])

    def test_observed_metadata_count_digest_and_actual_bytes_must_match(self):
        closure=self.verify()
        for changes in [dict(representation_digest=a.digest(b'other')),dict(bytes_written=3),
                        dict(representation_metadata={'media_type':'text/html','media_parameters':{},'languages':[]})]:
            with self.assertRaises(ValueError):a.correlate_delivery(self.delivery(**changes),closure)
        with self.assertRaises(ValueError):a.correlate_delivery(self.delivery(),closure,content_bytes=b'other')

    def test_coded_variant_and_decoded_representation_are_separate_checks(self):
        encoded=gzip.compress(self.body,mtime=0);v=self.value
        v['content_variants']=[dict(representation_id=v['representations'][0]['representation_id'],content_codings=['gzip'],content_length=len(encoded),content_digest=a.digest(encoded))]
        event=self.delivery(bytes_written=len(encoded),representation_digest=a.digest(encoded),content_codings=['gzip'],decoded_bytes=len(self.body),decoded_digest=a.digest(self.body))
        result=a.correlate_delivery(event,self.verify(),content_bytes=encoded,decoded_bytes=self.body)
        self.assertTrue(result['content_bytes_checked']);self.assertTrue(result['representation_bytes_match'])
        event['http']['content_digest']=a.digest(b'other')
        with self.assertRaises(a.AssetError):a.correlate_delivery(event,self.verify())

    def test_ranges_head_missing_metadata_and_upstream_hops_never_become_full_delivery(self):
        closure=self.verify()
        partial=self.delivery(status=206,bytes_written=3,observed_digest=a.digest(self.body[:3]),byte_range=dict(first=0,last=2,complete_length=len(self.body)))
        r=a.correlate_delivery(partial,closure,content_bytes=self.body[:3])
        self.assertEqual(r['binding'],'range_reference_consistent');self.assertFalse(r['representation_bytes_match'])
        for modify in [lambda e:e['http'].update(representation_metadata=None),lambda e:e['http'].update(hop_role='upstream')]:
            event=self.delivery();modify(event);self.assertFalse(a.correlate_delivery(event,closure)['representation_bytes_match'])

    def test_client_claim_and_unpublished_derivation_are_not_promoted(self):
        event=free_fixtures()[0]['report'];event['asset_ref']=ref(self.value)
        r=a.correlate_delivery(event,self.verify());self.assertEqual(r['binding'],'referenced_only')
        event['derived_from']=[ref(self.value)]
        with self.assertRaisesRegex(a.AssetError,'derivation differs'):a.correlate_delivery(event,self.verify())

    def test_derived_only_claim_cannot_reference_a_future_source(self):
        self.value['published_at']=later(1);self.clock=lambda:later(1)
        event=free_fixtures()[0]['report'];event.update(occurred_at=NOW,derived_from=[ref(self.value)])
        event['operation']['ended_at']=NOW
        with self.assertRaisesRegex(a.AssetError,'precedes source'):
            a.correlate_delivery(event,self.verify())

    def test_cross_origin_derivation_requires_both_independent_publishers(self):
        def replace(v):return json.loads(json.dumps(v).replace(ORIGIN,'https://second.example'))
        meta,pol,op=map(replace,asset_fixture())
        child=manifest(meta,pol,op,self.body,number=2,parents=[ref(self.value)])
        signed=[a.sign_manifest(a.raw(m),ec.derive_private_key(1,ec.SECP256R1()),m['publication']['key_id']) for m in (self.value,child)]
        docs={a.digest(a.raw(x)):a.raw(x) for x in (self.meta,self.pol,meta,pol)}
        data=a.raw(dict(protocol_version=a.VERSION,profile=a.PROFILE,roots=[ref(child)],manifests=signed,
            documents=[dict(digest=sha,body_b64url=crypto.b64u(body)) for sha,body in docs.items()]))
        trust=a.AssetTrust({ORIGIN:self.snap,meta['origin']:snapshot(meta,pol)},docs)
        result=self.verify(raw=data,trust=trust);self.assertEqual(len(result.manifests),2)
        self.assertFalse(result.resource_permission_granted)
        with self.assertRaisesRegex(a.AssetError,'independently current snapshot'):
            self.verify(raw=data,trust=a.AssetTrust({meta['origin']:snapshot(meta,pol)},docs))


def schema_cases():
    meta,pol,op=asset_fixture();m=manifest(meta,pol,op);b=a.load(bundle([m],meta,pol),a.MAX_BUNDLE)
    cases=[dict(name='native_asset_manifest',definition='manifest',value=m,structural_valid=True,semantic_valid=True),
           dict(name='native_asset_bundle',definition='bundle',value=b,structural_valid=True,semantic_valid=True)]
    def changed(name,definition,mutate,valid=False):
        value=copy.deepcopy(m if definition=='manifest' else b);mutate(value)
        cases.append(dict(name=name,definition=definition,value=value,structural_valid=valid,semantic_valid=False))
    changed('legacy_version','manifest',lambda v:v.update(protocol_version='1.2.0-draft.2'))
    changed('unknown_manifest_field','manifest',lambda v:v.update(extra=True))
    changed('missing_publication','manifest',lambda v:v.pop('publication'))
    changed('wrong_profile','bundle',lambda v:v.update(profile='assets'))
    changed('missing_closure','bundle',lambda v:v.update(manifests=[]))
    changed('duplicate_root','bundle',lambda v:v['roots'].append(v['roots'][0]))
    changed('invalid_jws','bundle',lambda v:v.update(manifests=['not-jws']))
    changed('bad_document_digest','bundle',lambda v:v['documents'][0].update(digest='wrong'))
    changed('negative_representation_size','manifest',lambda v:v['representations'][0].update(decoded_length=-1))
    changed('representation_extra_field','manifest',lambda v:v['representations'][0].update(extra=1))
    changed('origin_substitution_semantic','manifest',lambda v:v.update(origin='https://other.example'),True)
    changed('operation_endpoint_substitution_semantic','manifest',lambda v:v['publication'].update(endpoint=ORIGIN+'/versions/other'),True)
    changed('context_digest_substitution_semantic','manifest',lambda v:v.update(policy_digest='sha256:'+'0'*64),True)
    changed('lexical_fraction_semantic','manifest',lambda v:v['representations'][0].update(decoded_length=16.0),True)
    return dict(schema='asset',cases=cases)


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