from contextlib import closing
import sys,sqlite3,tempfile,json,uuid,unittest
from pathlib import Path
from datetime import datetime,timezone,timedelta
from unittest.mock import Mock
sys.path.insert(0,'/opt/charter-live/live_deploy_v2')
import order_dispatch_guard as g
from bybit_demo_client import BybitDemoClient
class GuardTests(unittest.TestCase):
 def setUp(self):
  self.tmp=tempfile.TemporaryDirectory();self.addCleanup(self.tmp.cleanup)
  g.DB_PATH=Path(self.tmp.name)/'queue.db';g.AUDIT_PATH=Path(self.tmp.name)/'audit.jsonl'
  self.rid=str(uuid.uuid4());self.now=datetime.now(timezone.utc)
  self.payload={'symbol':'NEARUSDT.P','side':'sell','sent_at':self.now.isoformat()}
  with closing(sqlite3.connect(g.DB_PATH)) as c, c:c.execute('create table queue(request_id text,status text,received_at text,payload text)')
  self.put()
  self.client=BybitDemoClient.__new__(BybitDemoClient);self.client.api_key='test';self.client.secret='test';self.client.session=Mock()
  self.client.session.post.return_value.json.return_value={'retCode':0,'result':{'orderId':'fake-order'}}
 def put(self,status='pending',payload=None,received=None):
  with closing(sqlite3.connect(g.DB_PATH)) as c, c:
   c.execute('delete from queue');c.execute('insert into queue values(?,?,?,?)',(self.rid,status,(received or self.now).isoformat(),json.dumps(payload or self.payload)))
 def send(self,**kw):
  args={'symbol':'NEARUSDT','side':'sell','qty':1,'order_link_id':'tv2-'+self.rid.replace('-','')};args.update(kw)
  return self.client.create_market_order(**args)
 def blocked(self,**kw):
  with self.assertRaises(RuntimeError):self.send(**kw)
  self.client.session.post.assert_not_called()
 def test_unlinked_order_blocked(self):self.blocked(order_link_id='')
 def test_unknown_request_blocked(self):self.blocked(order_link_id='tv2-'+uuid.uuid4().hex)
 def test_terminal_request_blocked(self):self.put(status='completed');self.blocked()
 def test_stale_alert_blocked(self):self.put(payload=dict(self.payload,sent_at=(self.now-timedelta(seconds=90)).isoformat()));self.blocked()
 def test_missing_time_blocked(self):self.put(payload={'symbol':'NEARUSDT.P','side':'sell'});self.blocked()
 def test_stale_queue_blocked(self):self.put(received=self.now-timedelta(seconds=700));self.blocked()
 def test_wrong_asset_and_side_blocked(self):
  self.blocked(symbol='BTCUSDT');self.blocked(side='buy')
 def test_fresh_linked_order_audited(self):
  self.send(reduce_only=True)
  sent=self.client.session.post.call_args.kwargs['json'];self.assertTrue(sent['reduceOnly']);self.assertEqual(len(sent['orderLinkId']),36)
  lines=[json.loads(x) for x in g.AUDIT_PATH.read_text().splitlines()];self.assertEqual([x['stage'] for x in lines],['attempt','response']);self.assertEqual(lines[-1]['orderId'],'fake-order')
 def test_all_assets(self):
  for symbol in ('BTCUSDT','AXSUSDT','1000BONKUSDT','ENAUSDT'):
   self.put(payload=dict(self.payload,symbol=symbol+'.P'));self.send(symbol=symbol)
 def test_unlinked_close_blocked(self):self.blocked(order_link_id='',reduce_only=True)
 def test_no_audit_no_network(self):
  g.AUDIT_PATH=Path(self.tmp.name)/'missing'/'audit.jsonl'
  with self.assertRaises(OSError):self.send()
  self.client.session.post.assert_not_called()
 def test_other_order_mutation_blocked(self):
  with self.assertRaises(RuntimeError):self.client._request('POST','/v5/order/create-batch',{})
  self.client.session.post.assert_not_called()
if __name__=='__main__':unittest.main()
