Skip to content

Commit 9b799f2

Browse files
AdamGlusteinclaude
andauthored
Support datetime/timedelta arithmetic between csp edges (#676)
Signed-off-by: Adam Glustein <adam.glustein@point72.com> Co-authored-by: Claude Opus 4.5 <noreply@anthropic.com>
1 parent 162f650 commit 9b799f2

2 files changed

Lines changed: 44 additions & 0 deletions

File tree

csp/math.py

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,5 @@
11
import math
2+
from datetime import datetime, timedelta
23
from functools import lru_cache
34
from typing import List, TypeVar, get_origin
45

@@ -270,6 +271,18 @@ def generic_type(x: ts["T"], y: ts["T"]) -> ts[generic_out_type]:
270271
if csp.valid(x, y):
271272
return op_lambda(x, y)
272273

274+
# Special case: datetime - datetime returns timedelta
275+
@_node_internal_use(name=name)
276+
def datetime_sub_type(x: ts[datetime], y: ts[datetime]) -> ts[timedelta]:
277+
if csp.valid(x, y):
278+
return op_lambda(x, y)
279+
280+
# Special case: datetime +/- timedelta returns datetime
281+
@_node_internal_use(name=name)
282+
def datetime_timedelta_type(x: ts[datetime], y: ts[timedelta]) -> ts[datetime]:
283+
if csp.valid(x, y):
284+
return op_lambda(x, y)
285+
273286
def comp(x: ts["T"], y: ts["U"]):
274287
if get_origin(x.tstype.typ) in [Numpy1DArray, NumpyNDArray] or get_origin(y.tstype.typ) in [
275288
Numpy1DArray,
@@ -280,6 +293,10 @@ def comp(x: ts["T"], y: ts["U"]):
280293
return float_type(x, y)
281294
elif x.tstype.typ is int and y.tstype.typ is int:
282295
return int_type(x, y)
296+
elif name == "sub" and x.tstype.typ is datetime and y.tstype.typ is datetime:
297+
return datetime_sub_type(x, y)
298+
elif name in ("add", "sub") and x.tstype.typ is datetime and y.tstype.typ is timedelta:
299+
return datetime_timedelta_type(x, y)
283300
return generic_type(x, y)
284301

285302
comp.__name__ = name

csp/tests/test_math.py

Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -281,6 +281,33 @@ def graph(use_promotion: bool):
281281
[v[1] for v in results[op.__name__ + "-rev"]], [comp(y, x) for x, y in zip(xv, yv)], op.__name__
282282
)
283283

284+
def test_arithmetic_time_ops(self):
285+
"""Test datetime/timedelta arithmetic operations between edges."""
286+
287+
@csp.graph
288+
def graph():
289+
dt1 = datetime(2020, 1, 1, 12, 0, 0)
290+
dt2 = datetime(2020, 1, 1, 10, 0, 0)
291+
td = timedelta(hours=2)
292+
dt1_edge = csp.const(dt1)
293+
dt2_edge = csp.const(dt2)
294+
td_edge = csp.const(td)
295+
296+
# datetime - datetime -> timedelta
297+
csp.add_graph_output("dt_sub_dt", dt1_edge - dt2_edge)
298+
299+
# datetime + timedelta -> datetime
300+
csp.add_graph_output("dt_plus_td", dt1_edge + td_edge)
301+
302+
# datetime - timedelta -> datetime
303+
csp.add_graph_output("dt_minus_td", dt1_edge - td_edge)
304+
305+
st = datetime(2020, 1, 1)
306+
results = csp.run(graph, starttime=st, endtime=st + timedelta(seconds=1))
307+
self.assertEqual(results["dt_sub_dt"][0][1], timedelta(hours=2))
308+
self.assertEqual(results["dt_plus_td"][0][1], datetime(2020, 1, 1, 14, 0, 0))
309+
self.assertEqual(results["dt_minus_td"][0][1], datetime(2020, 1, 1, 10, 0, 0))
310+
284311
def test_boolean_ops(self):
285312
def graph():
286313
x = csp.default(csp.curve(bool, [(timedelta(seconds=s), s % 2 == 0) for s in range(1, 20)]), False)

0 commit comments

Comments
 (0)