1from __future__ import annotations
2
3import importlib
4from dataclasses import dataclass
5from tempfile import NamedTemporaryFile
6from typing import Any
7
8from daggerml import Uri
9from daggerml.api import Dag, apply_codecs
10from daggerml.contrib.adapters import get_adapter
11from daggerml.contrib.s3 import S3Store
12
13
14class PandasDataFrameCodec:
15 def __init__(self, dataframe_type: type[Any]):
16 self._dataframe_type = dataframe_type
17
18 def can_encode(self, value: Any) -> bool:
19 return isinstance(value, self._dataframe_type)
20
21 def encode(self, value: Any, dag: Dag) -> Uri:
22 with NamedTemporaryFile(suffix=".parquet") as tmp:
23 value.to_parquet(tmp.name)
24 return S3Store().put(filepath=tmp.name, suffix=".parquet")
25
26
27class PolarsDataFrameCodec:
28 def __init__(self, dataframe_type: type[Any]):
29 self._dataframe_type = dataframe_type
30
31 def can_encode(self, value: Any) -> bool:
32 return isinstance(value, self._dataframe_type)
33
34 def encode(self, value: Any, dag: Dag) -> Uri:
35 with NamedTemporaryFile(suffix=".parquet") as tmp:
36 value.write_parquet(tmp.name)
37 return S3Store().put(filepath=tmp.name, suffix=".parquet")
38
39
40@dataclass(frozen=True)
41class DelayedRef:
42 name: str
43
44
45@dataclass(frozen=True)
46class DelayedLoad:
47 dagname: str
48 nodename: str | None = None
49
50
51@dataclass(frozen=True)
52class DelayedRunnable:
53 uri: str
54 adapter: str
55 sub: Any | "DelayedRunnable" | None
56 kwargs: dict[str, Any]
57
58
59class DelayedActionCodec:
60 def can_encode(self, value: Any) -> bool:
61 return isinstance(value, (DelayedRef, DelayedLoad, DelayedRunnable))
62
63 def encode(self, value: DelayedRef | DelayedLoad | DelayedRunnable, dag: "Dag"):
64 if isinstance(value, DelayedRef):
65 return apply_codecs(dag[value.name], dag=dag) # apply_codecs required for `Node` objects.
66 if isinstance(value, DelayedLoad):
67 return apply_codecs(dag.require(value.dagname, name=value.nodename), dag=dag) # apply_codecs required
68 adapter_spec = get_adapter(value.adapter)
69 # no need for `apply_codecs` because we return a Runnable which is recursed
70 return adapter_spec.resolve_runnable(value.uri, sub=value.sub, kwargs=value.kwargs)
71
72
73def _import_optional(module_name: str) -> Any | None:
74 try:
75 return importlib.import_module(module_name)
76 except ModuleNotFoundError as e:
77 if e.name == module_name:
78 return None
79 raise
80
81
82def literal_codecs() -> list[Any]:
83 codecs: list[Any] = [(1, DelayedActionCodec())]
84 pandas = _import_optional("pandas")
85 if pandas is not None:
86 codecs.append((1, PandasDataFrameCodec(pandas.DataFrame)))
87 polars = _import_optional("polars")
88 if polars is not None:
89 codecs.append((1, PolarsDataFrameCodec(polars.DataFrame)))
90 return codecs