Skip to content

Commit 042bcb3

Browse files
committed
Harden ABIDES execution scoring
1 parent 708af75 commit 042bcb3

1 file changed

Lines changed: 74 additions & 4 deletions

File tree

codeclash/arenas/abides/runtime/run_abides.py

Lines changed: 74 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -83,6 +83,68 @@ def make_player_agent(agent_class: type, player_name: str, agent_id: int):
8383
return guard_player_agent(agent)
8484

8585

86+
def is_recorded_execution(exchange: ExchangeAgent, order) -> bool:
87+
order_book = getattr(exchange, "order_books", {}).get(getattr(order, "symbol", None))
88+
if order_book is None:
89+
return False
90+
91+
order_id = getattr(order, "order_id", None)
92+
try:
93+
quantity = int(order.quantity)
94+
except (TypeError, ValueError, AttributeError):
95+
return False
96+
97+
for history_window in getattr(order_book, "history", []):
98+
order_record = history_window.get(order_id)
99+
if not order_record:
100+
continue
101+
for _timestamp, transaction_quantity in order_record.get("transactions", []):
102+
try:
103+
if int(transaction_quantity) == quantity:
104+
return True
105+
except (TypeError, ValueError):
106+
continue
107+
return False
108+
109+
110+
def instrument_player_orders(player_agents: dict[str, TradingAgent], exchange: ExchangeAgent) -> dict[str, set]:
111+
submitted_order_ids = {player: set() for player in player_agents}
112+
113+
for player, agent in player_agents.items():
114+
original_send_message = agent.sendMessage
115+
116+
def tracked_send_message(*args, _agent=agent, _player=player, _original=original_send_message, **kwargs):
117+
recipient_id = args[0] if args else kwargs.get("recipientID", kwargs.get("recipient_id"))
118+
msg = args[1] if len(args) > 1 else kwargs.get("msg")
119+
body = getattr(msg, "body", {})
120+
order = body.get("order") if body.get("msg") == "LIMIT_ORDER" else None
121+
if recipient_id == exchange.id and getattr(order, "agent_id", None) == _agent.id:
122+
submitted_order_ids[_player].add(getattr(order, "order_id", None))
123+
return _original(*args, **kwargs)
124+
125+
agent.sendMessage = tracked_send_message
126+
127+
return submitted_order_ids
128+
129+
130+
def instrument_order_books(exchange: ExchangeAgent) -> dict[str, int]:
131+
order_book_depth = {"count": 0}
132+
133+
for order_book in exchange.order_books.values():
134+
original_handle_limit_order = order_book.handleLimitOrder
135+
136+
def tracked_handle_limit_order(*args, _original=original_handle_limit_order, **kwargs):
137+
order_book_depth["count"] += 1
138+
try:
139+
return _original(*args, **kwargs)
140+
finally:
141+
order_book_depth["count"] -= 1
142+
143+
order_book.handleLimitOrder = tracked_handle_limit_order
144+
145+
return order_book_depth
146+
147+
86148
def guard_player_agent(agent: TradingAgent) -> TradingAgent:
87149
agent._codeclash_error = None
88150
agent._codeclash_traceback = None
@@ -203,6 +265,8 @@ def make_world_agents(agent_classes: dict[str, type], *, sim_idx: int, market_mi
203265
agent.log_to_file = False
204266

205267
ledgers = {player: {"CASH": STARTING_CASH, SYMBOL: 0} for player in player_agents}
268+
submitted_order_ids = instrument_player_orders(player_agents, exchange)
269+
order_book_depth = instrument_order_books(exchange)
206270
original_exchange_send = exchange.sendMessage
207271

208272
def scored_exchange_send(*args, **kwargs):
@@ -211,10 +275,16 @@ def scored_exchange_send(*args, **kwargs):
211275
player = player_ids.get(recipient_id)
212276
if player and getattr(msg, "body", {}).get("msg") == "ORDER_EXECUTED":
213277
order = msg.body["order"]
214-
quantity = int(order.quantity)
215-
signed_quantity = quantity if order.is_buy_order else -quantity
216-
ledgers[player][SYMBOL] += signed_quantity
217-
ledgers[player]["CASH"] -= signed_quantity * int(order.fill_price)
278+
order_id = getattr(order, "order_id", None)
279+
if (
280+
order_book_depth["count"] > 0
281+
and order_id in submitted_order_ids[player]
282+
and is_recorded_execution(exchange, order)
283+
):
284+
quantity = int(order.quantity)
285+
signed_quantity = quantity if order.is_buy_order else -quantity
286+
ledgers[player][SYMBOL] += signed_quantity
287+
ledgers[player]["CASH"] -= signed_quantity * int(order.fill_price)
218288
return original_exchange_send(*args, **kwargs)
219289

220290
exchange.sendMessage = scored_exchange_send

0 commit comments

Comments
 (0)