mrfkit 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- mrfkit/__init__.py +68 -0
- mrfkit/__main__.py +3 -0
- mrfkit/cli.py +79 -0
- mrfkit/codes.py +1607 -0
- mrfkit/csv_reader.py +324 -0
- mrfkit/files.py +651 -0
- mrfkit/headers.py +737 -0
- mrfkit/json_reader.py +603 -0
- mrfkit/payers.py +2755 -0
- mrfkit/records.py +260 -0
- mrfkit/reference.py +136 -0
- mrfkit/sinks.py +126 -0
- mrfkit/tabular.py +651 -0
- mrfkit/tic.py +660 -0
- mrfkit/values.py +627 -0
- mrfkit-0.1.0.dist-info/METADATA +136 -0
- mrfkit-0.1.0.dist-info/RECORD +21 -0
- mrfkit-0.1.0.dist-info/WHEEL +4 -0
- mrfkit-0.1.0.dist-info/entry_points.txt +2 -0
- mrfkit-0.1.0.dist-info/licenses/LICENSE +201 -0
- mrfkit-0.1.0.dist-info/licenses/NOTICE +4 -0
mrfkit/records.py
ADDED
|
@@ -0,0 +1,260 @@
|
|
|
1
|
+
"""The rows mrfkit produces.
|
|
2
|
+
|
|
3
|
+
Every record carries its charge item's key (code, code_type, billing_class,
|
|
4
|
+
setting, modifiers), so standard charges, payer rates and unmapped cells join
|
|
5
|
+
back to their item without a database. Values are what the file says after
|
|
6
|
+
normalization: mrfkit never truncates text or clamps numbers.
|
|
7
|
+
|
|
8
|
+
Insurer Transparency in Coverage records (``Tic*``) come last and use the TiC
|
|
9
|
+
schema's own field names.
|
|
10
|
+
"""
|
|
11
|
+
|
|
12
|
+
from __future__ import annotations
|
|
13
|
+
|
|
14
|
+
from dataclasses import dataclass
|
|
15
|
+
from typing import ClassVar, Optional
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
@dataclass(slots=True)
|
|
19
|
+
class ChargeItem:
|
|
20
|
+
"""One billable item. Emitted once per distinct item key in a file."""
|
|
21
|
+
|
|
22
|
+
TABLE: ClassVar[str] = "charge_items"
|
|
23
|
+
|
|
24
|
+
code: Optional[str]
|
|
25
|
+
code_type: Optional[str]
|
|
26
|
+
description: Optional[str]
|
|
27
|
+
billing_class: str = ""
|
|
28
|
+
setting: str = ""
|
|
29
|
+
modifiers: Optional[str] = None
|
|
30
|
+
drug_unit_of_measurement: Optional[str] = None
|
|
31
|
+
drug_type_of_measurement: Optional[str] = None
|
|
32
|
+
additional_generic_notes: Optional[str] = None
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
@dataclass(slots=True)
|
|
36
|
+
class StandardCharge:
|
|
37
|
+
"""The payer-independent prices for an item."""
|
|
38
|
+
|
|
39
|
+
TABLE: ClassVar[str] = "standard_charges"
|
|
40
|
+
|
|
41
|
+
code: Optional[str]
|
|
42
|
+
code_type: Optional[str]
|
|
43
|
+
description: Optional[str]
|
|
44
|
+
billing_class: str = ""
|
|
45
|
+
setting: str = ""
|
|
46
|
+
modifiers: Optional[str] = None
|
|
47
|
+
gross_charge: Optional[float] = None
|
|
48
|
+
discounted_cash_price: Optional[float] = None
|
|
49
|
+
min_negotiated_rate: Optional[float] = None
|
|
50
|
+
max_negotiated_rate: Optional[float] = None
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
@dataclass(slots=True)
|
|
54
|
+
class PayerRate:
|
|
55
|
+
"""One negotiated rate for an item under one payer and plan.
|
|
56
|
+
|
|
57
|
+
``billing_class``, ``setting`` and ``modifiers`` are the item's key.
|
|
58
|
+
``rate_billing_class`` and ``rate_setting`` are what the rate itself says,
|
|
59
|
+
which can differ from the item's.
|
|
60
|
+
"""
|
|
61
|
+
|
|
62
|
+
TABLE: ClassVar[str] = "payer_rates"
|
|
63
|
+
|
|
64
|
+
code: Optional[str]
|
|
65
|
+
code_type: Optional[str]
|
|
66
|
+
description: Optional[str]
|
|
67
|
+
billing_class: str = ""
|
|
68
|
+
setting: str = ""
|
|
69
|
+
modifiers: Optional[str] = None
|
|
70
|
+
payer_name: Optional[str] = None # normalized
|
|
71
|
+
raw_payer_name: Optional[str] = None # as written in the file
|
|
72
|
+
plan_name: Optional[str] = None
|
|
73
|
+
plan_category: str = "Other"
|
|
74
|
+
plan_network: Optional[str] = None
|
|
75
|
+
negotiated_rate: Optional[float] = None
|
|
76
|
+
negotiated_percentage: Optional[float] = None
|
|
77
|
+
negotiated_algorithm: Optional[str] = None
|
|
78
|
+
methodology: Optional[str] = None # as written, minus audit notes and URLs
|
|
79
|
+
methodology_type: Optional[str] = None # normalized category
|
|
80
|
+
estimated_amount: Optional[float] = None
|
|
81
|
+
rate_billing_class: str = ""
|
|
82
|
+
rate_setting: str = ""
|
|
83
|
+
additional_notes: Optional[str] = None
|
|
84
|
+
footnote: Optional[str] = None
|
|
85
|
+
median_amount: Optional[float] = None
|
|
86
|
+
pct_10: Optional[float] = None
|
|
87
|
+
pct_90: Optional[float] = None
|
|
88
|
+
claim_count: Optional[int] = None
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
@dataclass(slots=True)
|
|
92
|
+
class UnmappedCell:
|
|
93
|
+
"""A value from a column mrfkit could not map to a field.
|
|
94
|
+
|
|
95
|
+
Kept so nothing in the file is silently lost, and so new header synonyms
|
|
96
|
+
can be found. Capped per column by the readers.
|
|
97
|
+
"""
|
|
98
|
+
|
|
99
|
+
TABLE: ClassVar[str] = "unmapped_cells"
|
|
100
|
+
|
|
101
|
+
code: Optional[str]
|
|
102
|
+
code_type: Optional[str]
|
|
103
|
+
description: Optional[str]
|
|
104
|
+
billing_class: str = ""
|
|
105
|
+
setting: str = ""
|
|
106
|
+
modifiers: Optional[str] = None
|
|
107
|
+
source_column: str = ""
|
|
108
|
+
value: Optional[str] = None
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
@dataclass(slots=True)
|
|
112
|
+
class FileMetadata:
|
|
113
|
+
"""What the file says about itself: hospital, license, dates, attestation.
|
|
114
|
+
|
|
115
|
+
Several locations or addresses are joined with ``|``, as the CMS CSV
|
|
116
|
+
template writes them.
|
|
117
|
+
"""
|
|
118
|
+
|
|
119
|
+
TABLE: ClassVar[str] = "file_metadata"
|
|
120
|
+
|
|
121
|
+
hospital_name: Optional[str] = None
|
|
122
|
+
hospital_location: Optional[str] = None
|
|
123
|
+
hospital_address: Optional[str] = None
|
|
124
|
+
cms_certification_number: Optional[str] = None
|
|
125
|
+
license_number: Optional[str] = None
|
|
126
|
+
license_state: Optional[str] = None
|
|
127
|
+
type_2_npi: Optional[str] = None
|
|
128
|
+
ein: Optional[str] = None
|
|
129
|
+
last_updated_on: Optional[str] = None
|
|
130
|
+
version: Optional[str] = None
|
|
131
|
+
attestation: Optional[str] = None
|
|
132
|
+
attester_name: Optional[str] = None
|
|
133
|
+
confirm_attestation: Optional[bool] = None
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
@dataclass(slots=True)
|
|
137
|
+
class HeaderMapping:
|
|
138
|
+
"""How one source column (or JSON key) was read.
|
|
139
|
+
|
|
140
|
+
``mapped_to`` is a canonical field name, a layout tag such as
|
|
141
|
+
``wide_payer:negotiated_rate``, or None for an unmapped column.
|
|
142
|
+
"""
|
|
143
|
+
|
|
144
|
+
TABLE: ClassVar[str] = "header_mappings"
|
|
145
|
+
|
|
146
|
+
source_header: str
|
|
147
|
+
normalized: str
|
|
148
|
+
mapped_to: Optional[str] = None
|
|
149
|
+
|
|
150
|
+
|
|
151
|
+
@dataclass(slots=True)
|
|
152
|
+
class ModifierInfo:
|
|
153
|
+
"""A modifier the file defines (CMS 3.0 ``modifier_information``).
|
|
154
|
+
|
|
155
|
+
One general record per modifier code, then one per payer with its own
|
|
156
|
+
description. ``payer_name`` is normalized; ``raw_payer_name`` is as written.
|
|
157
|
+
"""
|
|
158
|
+
|
|
159
|
+
TABLE: ClassVar[str] = "modifiers"
|
|
160
|
+
|
|
161
|
+
code: str
|
|
162
|
+
description: Optional[str] = None
|
|
163
|
+
setting: Optional[str] = None
|
|
164
|
+
payer_name: Optional[str] = None
|
|
165
|
+
raw_payer_name: Optional[str] = None
|
|
166
|
+
plan_name: Optional[str] = None
|
|
167
|
+
payer_description: Optional[str] = None
|
|
168
|
+
|
|
169
|
+
|
|
170
|
+
# ---------------------------------------------------------------------------
|
|
171
|
+
# Insurer Transparency in Coverage (TiC) files
|
|
172
|
+
# ---------------------------------------------------------------------------
|
|
173
|
+
|
|
174
|
+
@dataclass(slots=True)
|
|
175
|
+
class TicFileMetadata:
|
|
176
|
+
"""What a TiC in-network rates file says about itself.
|
|
177
|
+
|
|
178
|
+
``network_name`` is the first provider reference's network label, which
|
|
179
|
+
payers repeat on every reference.
|
|
180
|
+
"""
|
|
181
|
+
|
|
182
|
+
TABLE: ClassVar[str] = "tic_file_metadata"
|
|
183
|
+
|
|
184
|
+
reporting_entity_name: Optional[str] = None
|
|
185
|
+
reporting_entity_type: Optional[str] = None
|
|
186
|
+
last_updated_on: Optional[str] = None
|
|
187
|
+
version: Optional[str] = None
|
|
188
|
+
network_name: Optional[str] = None
|
|
189
|
+
|
|
190
|
+
|
|
191
|
+
@dataclass(slots=True)
|
|
192
|
+
class TicProviderGroup:
|
|
193
|
+
"""One provider group, a TIN and its NPIs, behind a provider reference.
|
|
194
|
+
|
|
195
|
+
``group_id`` is the file's own ``provider_group_id``, so it is unique
|
|
196
|
+
only within one file, and a reference holding several groups has several
|
|
197
|
+
rows. A group listed inline on a rate gets a synthetic ``inline-N`` id.
|
|
198
|
+
``tin`` is digits only. Several NPIs are joined with ``|``.
|
|
199
|
+
"""
|
|
200
|
+
|
|
201
|
+
TABLE: ClassVar[str] = "tic_provider_groups"
|
|
202
|
+
|
|
203
|
+
group_id: str
|
|
204
|
+
tin_type: Optional[str]
|
|
205
|
+
tin: str
|
|
206
|
+
npis: str
|
|
207
|
+
business_name: Optional[str] = None
|
|
208
|
+
|
|
209
|
+
|
|
210
|
+
@dataclass(slots=True)
|
|
211
|
+
class TicRate:
|
|
212
|
+
"""One negotiated price for a billing code, shared by provider groups.
|
|
213
|
+
|
|
214
|
+
Identical prices for one code are merged into a single record whose
|
|
215
|
+
``provider_group_ids`` (``|``-joined, sorted) lists every
|
|
216
|
+
:class:`TicProviderGroup` that has it, instead of one row per provider
|
|
217
|
+
or NPI. A ``percentage`` price fills ``negotiated_percentage``, every
|
|
218
|
+
other type ``negotiated_rate``. ``expiration_date`` is ``9999-12-31``
|
|
219
|
+
when the price does not expire, and ``setting`` is empty when the file
|
|
220
|
+
gives none. Several service codes or modifiers are joined with ``|``.
|
|
221
|
+
"""
|
|
222
|
+
|
|
223
|
+
TABLE: ClassVar[str] = "tic_rates"
|
|
224
|
+
|
|
225
|
+
billing_code: Optional[str]
|
|
226
|
+
billing_code_type: Optional[str]
|
|
227
|
+
billing_code_type_version: Optional[str] = None
|
|
228
|
+
negotiation_arrangement: Optional[str] = None
|
|
229
|
+
negotiated_type: str = ""
|
|
230
|
+
negotiated_rate: Optional[float] = None
|
|
231
|
+
negotiated_percentage: Optional[float] = None
|
|
232
|
+
expiration_date: str = "9999-12-31"
|
|
233
|
+
billing_class: str = ""
|
|
234
|
+
setting: str = ""
|
|
235
|
+
service_codes: Optional[str] = None
|
|
236
|
+
modifiers: Optional[str] = None
|
|
237
|
+
provider_group_ids: str = ""
|
|
238
|
+
|
|
239
|
+
|
|
240
|
+
@dataclass(slots=True)
|
|
241
|
+
class TicIndexEntry:
|
|
242
|
+
"""One in-network file a plan uses, from a TiC table of contents.
|
|
243
|
+
|
|
244
|
+
One record per (reporting plan, in-network file) pair of a
|
|
245
|
+
``reporting_structure``: a structure with two plans and three files
|
|
246
|
+
gives six records. ``allowed_amount_location`` is the structure's
|
|
247
|
+
allowed-amount file, which mrfkit does not read.
|
|
248
|
+
"""
|
|
249
|
+
|
|
250
|
+
TABLE: ClassVar[str] = "tic_index"
|
|
251
|
+
|
|
252
|
+
reporting_entity_name: Optional[str] = None
|
|
253
|
+
reporting_entity_type: Optional[str] = None
|
|
254
|
+
plan_name: Optional[str] = None
|
|
255
|
+
plan_id: Optional[str] = None
|
|
256
|
+
plan_id_type: Optional[str] = None
|
|
257
|
+
plan_market_type: Optional[str] = None
|
|
258
|
+
in_network_location: str = ""
|
|
259
|
+
in_network_description: Optional[str] = None
|
|
260
|
+
allowed_amount_location: Optional[str] = None
|
mrfkit/reference.py
ADDED
|
@@ -0,0 +1,136 @@
|
|
|
1
|
+
"""Optional reference data and parse counters for the normalizers.
|
|
2
|
+
|
|
3
|
+
Every normalizer works with no reference data at all. When you have lookup
|
|
4
|
+
tables (published code lists, payer aliases you curate, payers you already
|
|
5
|
+
know), put them in a ``ReferenceData`` and pass it as ``ref=``. When you want
|
|
6
|
+
to know how many rows were rejected or repaired, pass a ``ParseStats`` as
|
|
7
|
+
``stats=``.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from __future__ import annotations
|
|
11
|
+
|
|
12
|
+
import logging
|
|
13
|
+
from dataclasses import dataclass, field
|
|
14
|
+
from typing import Dict, FrozenSet, List, Mapping, Optional
|
|
15
|
+
|
|
16
|
+
from mrfkit.payers import _clean_payer_key
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
@dataclass(frozen=True)
|
|
20
|
+
class ReferenceData:
|
|
21
|
+
"""Lookup tables that sharpen normalization. Every field is optional.
|
|
22
|
+
|
|
23
|
+
code_prefixes
|
|
24
|
+
``{'CPT': frozenset(codes), 'HCPCS': frozenset(codes)}`` of published
|
|
25
|
+
codes. When set, a modifier baked into a code (``73721TC``) is split
|
|
26
|
+
only if the remaining code is in this list. ``None`` skips the check;
|
|
27
|
+
an empty set for a type rejects every split of that type.
|
|
28
|
+
modifier_validity
|
|
29
|
+
CPT code -> frozenset of digit modifiers CMS marks valid for it
|
|
30
|
+
(``{'73721': frozenset({'50'})}``). Digit modifiers baked into a code
|
|
31
|
+
split only when this confirms the pair. ``None`` means they never
|
|
32
|
+
split.
|
|
33
|
+
code_descriptions
|
|
34
|
+
CPT code -> a short description you hold the rights to. When set, a
|
|
35
|
+
digit modifier split also requires the row description to resemble
|
|
36
|
+
it.
|
|
37
|
+
payer_aliases
|
|
38
|
+
Raw payer name -> canonical payer name, checked before the built-in
|
|
39
|
+
map. Keys are cleaned (quotes, brackets, underscores, spacing) and
|
|
40
|
+
uppercased for you.
|
|
41
|
+
payer_match_index
|
|
42
|
+
``payer_match_key(name)`` -> canonical payer name of payers you
|
|
43
|
+
already know. Spelling variants of an otherwise unknown payer fold
|
|
44
|
+
onto the name here. When the index is not empty, new payers are added
|
|
45
|
+
to it as they are seen, so later variants in the same run fold onto
|
|
46
|
+
the first spelling.
|
|
47
|
+
methodology_aliases
|
|
48
|
+
Raw methodology text -> methodology type, checked before the built-in
|
|
49
|
+
map. Keys are uppercased for you.
|
|
50
|
+
"""
|
|
51
|
+
|
|
52
|
+
code_prefixes: Optional[Mapping[str, FrozenSet[str]]] = None
|
|
53
|
+
modifier_validity: Optional[Mapping[str, FrozenSet[str]]] = None
|
|
54
|
+
code_descriptions: Optional[Mapping[str, str]] = None
|
|
55
|
+
payer_aliases: Dict[str, str] = field(default_factory=dict)
|
|
56
|
+
payer_match_index: Dict[str, str] = field(default_factory=dict)
|
|
57
|
+
methodology_aliases: Dict[str, str] = field(default_factory=dict)
|
|
58
|
+
|
|
59
|
+
def __post_init__(self) -> None:
|
|
60
|
+
aliases = {}
|
|
61
|
+
for alias, canonical in self.payer_aliases.items():
|
|
62
|
+
key = _clean_payer_key(alias)
|
|
63
|
+
if key and canonical:
|
|
64
|
+
aliases[key] = canonical
|
|
65
|
+
object.__setattr__(self, 'payer_aliases', aliases)
|
|
66
|
+
object.__setattr__(self, 'payer_match_index', dict(self.payer_match_index))
|
|
67
|
+
object.__setattr__(
|
|
68
|
+
self, 'methodology_aliases',
|
|
69
|
+
{k.upper(): v for k, v in self.methodology_aliases.items()},
|
|
70
|
+
)
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
@dataclass
|
|
74
|
+
class ParseStats:
|
|
75
|
+
"""Counts of rows the normalizers rejected or repaired.
|
|
76
|
+
|
|
77
|
+
Pass one instance per file as ``stats=`` to ``is_rejected_code`` and
|
|
78
|
+
``apply_baked_modifier_split``. Sample lists keep at most
|
|
79
|
+
``sample_limit`` entries, with each value cut short, so a file with
|
|
80
|
+
millions of bad rows cannot blow up memory. The TiC readers count each
|
|
81
|
+
JSON path they could not map in ``unmapped_paths``.
|
|
82
|
+
"""
|
|
83
|
+
|
|
84
|
+
sample_limit: int = 20
|
|
85
|
+
rows_read: int = 0
|
|
86
|
+
rows_skipped: int = 0
|
|
87
|
+
warnings: List[str] = field(default_factory=list)
|
|
88
|
+
rejected_codes: int = 0
|
|
89
|
+
rejected_code_samples: List[Dict[str, str]] = field(default_factory=list)
|
|
90
|
+
baked_modifier_splits: int = 0
|
|
91
|
+
baked_modifiers: Dict[str, int] = field(default_factory=dict)
|
|
92
|
+
baked_modifier_samples: List[Dict[str, str]] = field(default_factory=list)
|
|
93
|
+
unmapped_path_limit: int = 500
|
|
94
|
+
unmapped_paths: Dict[str, int] = field(default_factory=dict)
|
|
95
|
+
|
|
96
|
+
def warn(self, message: str) -> None:
|
|
97
|
+
"""Keep *message* for the caller and log it."""
|
|
98
|
+
self.warnings.append(message)
|
|
99
|
+
logging.getLogger("mrfkit").warning(message)
|
|
100
|
+
|
|
101
|
+
def record_unmapped_path(self, path: str) -> None:
|
|
102
|
+
"""Count one occurrence of *path*: a JSON path a reader could not map.
|
|
103
|
+
|
|
104
|
+
Only the path and a count are kept, never the value found there, so
|
|
105
|
+
free text in an unknown field cannot end up in the output. At most
|
|
106
|
+
``unmapped_path_limit`` distinct paths are kept.
|
|
107
|
+
"""
|
|
108
|
+
if path in self.unmapped_paths:
|
|
109
|
+
self.unmapped_paths[path] += 1
|
|
110
|
+
elif len(self.unmapped_paths) < self.unmapped_path_limit:
|
|
111
|
+
self.unmapped_paths[path] = 1
|
|
112
|
+
|
|
113
|
+
def record_rejected_code(
|
|
114
|
+
self, code: Optional[str], code_type: Optional[str],
|
|
115
|
+
) -> None:
|
|
116
|
+
self.rejected_codes += 1
|
|
117
|
+
if len(self.rejected_code_samples) < self.sample_limit:
|
|
118
|
+
self.rejected_code_samples.append({
|
|
119
|
+
'code': (code or '')[:32],
|
|
120
|
+
'code_type': (code_type or '')[:32],
|
|
121
|
+
})
|
|
122
|
+
|
|
123
|
+
def record_baked_modifier_split(
|
|
124
|
+
self, baked_modifier: Optional[str], source_code_type: Optional[str],
|
|
125
|
+
) -> None:
|
|
126
|
+
if not baked_modifier:
|
|
127
|
+
return
|
|
128
|
+
self.baked_modifier_splits += 1
|
|
129
|
+
self.baked_modifiers[baked_modifier] = (
|
|
130
|
+
self.baked_modifiers.get(baked_modifier, 0) + 1
|
|
131
|
+
)
|
|
132
|
+
if len(self.baked_modifier_samples) < self.sample_limit:
|
|
133
|
+
self.baked_modifier_samples.append({
|
|
134
|
+
'baked_modifier': baked_modifier[:8],
|
|
135
|
+
'source_code_type': (source_code_type or '')[:16],
|
|
136
|
+
})
|
mrfkit/sinks.py
ADDED
|
@@ -0,0 +1,126 @@
|
|
|
1
|
+
"""Write records to disk: one file per record type, in CSV or Parquet."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import csv
|
|
6
|
+
import dataclasses
|
|
7
|
+
import operator
|
|
8
|
+
import typing
|
|
9
|
+
from pathlib import Path
|
|
10
|
+
from typing import Any, Callable, Dict, List, Tuple, Type
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def _columns(record_type: type) -> Tuple[List[str], Callable[[Any], tuple]]:
|
|
14
|
+
names = [f.name for f in dataclasses.fields(record_type)]
|
|
15
|
+
getter = operator.attrgetter(*names)
|
|
16
|
+
if len(names) == 1:
|
|
17
|
+
return names, lambda r: (getter(r),)
|
|
18
|
+
return names, getter
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class CsvSink:
|
|
22
|
+
"""Write each record type to ``<out_dir>/<TABLE>.csv``. None becomes an empty cell."""
|
|
23
|
+
|
|
24
|
+
def __init__(self, out_dir: Path):
|
|
25
|
+
self.out_dir = Path(out_dir)
|
|
26
|
+
self.out_dir.mkdir(parents=True, exist_ok=True)
|
|
27
|
+
self._open: Dict[type, Tuple[Any, Any, Callable]] = {}
|
|
28
|
+
|
|
29
|
+
def write(self, record: Any) -> None:
|
|
30
|
+
entry = self._open.get(type(record))
|
|
31
|
+
if entry is None:
|
|
32
|
+
names, getter = _columns(type(record))
|
|
33
|
+
fh = open(self.out_dir / f"{record.TABLE}.csv", "w", newline="", encoding="utf-8")
|
|
34
|
+
writer = csv.writer(fh)
|
|
35
|
+
writer.writerow(names)
|
|
36
|
+
entry = self._open[type(record)] = (fh, writer, getter)
|
|
37
|
+
entry[1].writerow(entry[2](record))
|
|
38
|
+
|
|
39
|
+
def close(self) -> None:
|
|
40
|
+
for fh, _, _ in self._open.values():
|
|
41
|
+
fh.close()
|
|
42
|
+
self._open.clear()
|
|
43
|
+
|
|
44
|
+
def __enter__(self) -> "CsvSink":
|
|
45
|
+
return self
|
|
46
|
+
|
|
47
|
+
def __exit__(self, *exc) -> None:
|
|
48
|
+
self.close()
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
class ParquetSink:
|
|
52
|
+
"""Write each record type to ``<out_dir>/<TABLE>.parquet``, in row groups of *batch_size*."""
|
|
53
|
+
|
|
54
|
+
def __init__(self, out_dir: Path, batch_size: int = 50_000):
|
|
55
|
+
try:
|
|
56
|
+
import pyarrow
|
|
57
|
+
import pyarrow.parquet
|
|
58
|
+
except ImportError as exc:
|
|
59
|
+
raise ImportError('Parquet output needs pyarrow: pip install "mrfkit[parquet]"') from exc
|
|
60
|
+
self._pa = pyarrow
|
|
61
|
+
self._pq = pyarrow.parquet
|
|
62
|
+
self.out_dir = Path(out_dir)
|
|
63
|
+
self.out_dir.mkdir(parents=True, exist_ok=True)
|
|
64
|
+
self.batch_size = batch_size
|
|
65
|
+
self._rows: Dict[type, list] = {}
|
|
66
|
+
self._writers: Dict[type, Any] = {}
|
|
67
|
+
self._layout: Dict[type, Tuple[List[str], Callable, Any]] = {}
|
|
68
|
+
|
|
69
|
+
def write(self, record: Any) -> None:
|
|
70
|
+
record_type = type(record)
|
|
71
|
+
if record_type not in self._layout:
|
|
72
|
+
names, getter = _columns(record_type)
|
|
73
|
+
self._layout[record_type] = (names, getter, self._schema(record_type, names))
|
|
74
|
+
self._rows[record_type] = []
|
|
75
|
+
rows = self._rows[record_type]
|
|
76
|
+
rows.append(self._layout[record_type][1](record))
|
|
77
|
+
if len(rows) >= self.batch_size:
|
|
78
|
+
self._flush(record_type)
|
|
79
|
+
|
|
80
|
+
def _schema(self, record_type: type, names: List[str]):
|
|
81
|
+
pa = self._pa
|
|
82
|
+
hints = typing.get_type_hints(record_type)
|
|
83
|
+
by_type = {str: pa.string(), float: pa.float64(), int: pa.int64(), bool: pa.bool_()}
|
|
84
|
+
fields = []
|
|
85
|
+
for name in names:
|
|
86
|
+
args = [a for a in typing.get_args(hints[name]) if a is not type(None)] or [hints[name]]
|
|
87
|
+
fields.append(pa.field(name, by_type[args[0]]))
|
|
88
|
+
return pa.schema(fields)
|
|
89
|
+
|
|
90
|
+
def _flush(self, record_type: Type) -> None:
|
|
91
|
+
rows = self._rows[record_type]
|
|
92
|
+
if not rows:
|
|
93
|
+
return
|
|
94
|
+
names, _, schema = self._layout[record_type]
|
|
95
|
+
table = self._pa.Table.from_arrays(
|
|
96
|
+
[self._pa.array(col, type=schema.field(i).type) for i, col in enumerate(zip(*rows, strict=True))],
|
|
97
|
+
schema=schema,
|
|
98
|
+
)
|
|
99
|
+
writer = self._writers.get(record_type)
|
|
100
|
+
if writer is None:
|
|
101
|
+
path = self.out_dir / f"{record_type.TABLE}.parquet"
|
|
102
|
+
writer = self._writers[record_type] = self._pq.ParquetWriter(path, schema)
|
|
103
|
+
writer.write_table(table)
|
|
104
|
+
rows.clear()
|
|
105
|
+
|
|
106
|
+
def close(self) -> None:
|
|
107
|
+
for record_type in list(self._rows):
|
|
108
|
+
self._flush(record_type)
|
|
109
|
+
for writer in self._writers.values():
|
|
110
|
+
writer.close()
|
|
111
|
+
self._writers.clear()
|
|
112
|
+
|
|
113
|
+
def __enter__(self) -> "ParquetSink":
|
|
114
|
+
return self
|
|
115
|
+
|
|
116
|
+
def __exit__(self, *exc) -> None:
|
|
117
|
+
self.close()
|
|
118
|
+
|
|
119
|
+
|
|
120
|
+
def open_sink(out_dir: Path, fmt: str = "csv"):
|
|
121
|
+
"""Return a CSV or Parquet sink for *out_dir*."""
|
|
122
|
+
if fmt == "csv":
|
|
123
|
+
return CsvSink(out_dir)
|
|
124
|
+
if fmt == "parquet":
|
|
125
|
+
return ParquetSink(out_dir)
|
|
126
|
+
raise ValueError(f"Unknown output format {fmt!r}: use 'csv' or 'parquet'")
|