openbb-us-eia 1.0.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.
- openbb_us_eia/__init__.py +25 -0
- openbb_us_eia/models/__init__.py +1 -0
- openbb_us_eia/models/petroleum_status_report.py +241 -0
- openbb_us_eia/models/short_term_energy_outlook.py +284 -0
- openbb_us_eia/utils/__init__.py +1 -0
- openbb_us_eia/utils/constants.py +1272 -0
- openbb_us_eia/utils/helpers.py +43 -0
- openbb_us_eia/utils/py.typed +0 -0
- openbb_us_eia/utils/total_energy_map.json +982 -0
- openbb_us_eia-1.0.0.dist-info/METADATA +140 -0
- openbb_us_eia-1.0.0.dist-info/RECORD +13 -0
- openbb_us_eia-1.0.0.dist-info/WHEEL +4 -0
- openbb_us_eia-1.0.0.dist-info/entry_points.txt +3 -0
|
@@ -0,0 +1,25 @@
|
|
|
1
|
+
"""OpenBB EIA Provider Module."""
|
|
2
|
+
|
|
3
|
+
from openbb_core.provider.abstract.provider import Provider
|
|
4
|
+
from openbb_us_eia.models.petroleum_status_report import EiaPetroleumStatusReportFetcher
|
|
5
|
+
from openbb_us_eia.models.short_term_energy_outlook import (
|
|
6
|
+
EiaShortTermEnergyOutlookFetcher,
|
|
7
|
+
)
|
|
8
|
+
|
|
9
|
+
eia_provider = Provider(
|
|
10
|
+
name="eia",
|
|
11
|
+
website="https://eia.gov/",
|
|
12
|
+
description="The U.S. Energy Information Administration is committed to its free and open data"
|
|
13
|
+
+ " by making it available through an Application Programming Interface (API) and its open data tools."
|
|
14
|
+
+ " See https://www.eia.gov/opendata/ for more information.",
|
|
15
|
+
credentials=[
|
|
16
|
+
"api_key"
|
|
17
|
+
], # This is not required for the Weekly Petroleum Status Report
|
|
18
|
+
fetcher_dict={
|
|
19
|
+
"PetroleumStatusReport": EiaPetroleumStatusReportFetcher,
|
|
20
|
+
"ShortTermEnergyOutlook": EiaShortTermEnergyOutlookFetcher,
|
|
21
|
+
},
|
|
22
|
+
repr_name="U.S. Energy Information Administration (EIA) Open Data and API",
|
|
23
|
+
instructions="""Credentials are required for functions calling the EIA's API.
|
|
24
|
+
Register for a free key here: https://www.eia.gov/opendata/register.php""",
|
|
25
|
+
)
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
"""OpenBB EIA Provider Models."""
|
|
@@ -0,0 +1,241 @@
|
|
|
1
|
+
"""EIA Weekly Petroleum Status Report model."""
|
|
2
|
+
|
|
3
|
+
# pylint: disable=unused-argument
|
|
4
|
+
|
|
5
|
+
from typing import Any, Optional
|
|
6
|
+
|
|
7
|
+
from openbb_core.app.model.abstract.error import OpenBBError
|
|
8
|
+
from openbb_core.provider.abstract.fetcher import Fetcher
|
|
9
|
+
from openbb_core.provider.standard_models.petroleum_status_report import (
|
|
10
|
+
PetroleumStatusReportData,
|
|
11
|
+
PetroleumStatusReportQueryParams,
|
|
12
|
+
)
|
|
13
|
+
from openbb_core.provider.utils.errors import EmptyDataError
|
|
14
|
+
from openbb_us_eia.utils.constants import (
|
|
15
|
+
WpsrCategoryChoices,
|
|
16
|
+
WpsrCategoryType,
|
|
17
|
+
WpsrFileMap,
|
|
18
|
+
WpsrTableChoices,
|
|
19
|
+
WpsrTableMap,
|
|
20
|
+
)
|
|
21
|
+
from pydantic import Field
|
|
22
|
+
|
|
23
|
+
WpsrTableChoicesString = "\n ".join(WpsrTableChoices)
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
class EiaPetroleumStatusReportQueryParams(PetroleumStatusReportQueryParams):
|
|
27
|
+
"""EIA Petroleum Status Report Query Parameters.
|
|
28
|
+
|
|
29
|
+
Source: https://www.eia.gov/petroleum/supply/weekly/
|
|
30
|
+
"""
|
|
31
|
+
|
|
32
|
+
__json_schema_extra__ = {
|
|
33
|
+
"category": {
|
|
34
|
+
"multiiple_items_allowed": False,
|
|
35
|
+
"choices": WpsrCategoryChoices,
|
|
36
|
+
},
|
|
37
|
+
"table": {
|
|
38
|
+
"multiple_items_allowed": True,
|
|
39
|
+
"choices": WpsrTableChoices,
|
|
40
|
+
},
|
|
41
|
+
}
|
|
42
|
+
|
|
43
|
+
category: WpsrCategoryType = Field(
|
|
44
|
+
default="balance_sheet",
|
|
45
|
+
description="The group of data to be returned. The default is the balance sheet.",
|
|
46
|
+
)
|
|
47
|
+
table: Optional[str] = Field(
|
|
48
|
+
default=None,
|
|
49
|
+
description="The specific table element within the category to be returned,"
|
|
50
|
+
+ " default is 'stocks', if the category is 'weekly_estimates', else 'all'."
|
|
51
|
+
+ "\n Note: Choices represent all available tables from the entire collection and are not all"
|
|
52
|
+
+ " available for every category."
|
|
53
|
+
+ "\n Invalid choices will raise a ValidationError with a message"
|
|
54
|
+
+ " indicating the valid choices for the selected category."
|
|
55
|
+
+ "\n Choices are:"
|
|
56
|
+
+ f"\n {WpsrTableChoicesString}\n ",
|
|
57
|
+
)
|
|
58
|
+
use_cache: bool = Field(
|
|
59
|
+
default=True,
|
|
60
|
+
description="Subsequent requests for the same source data are cached for the session using ALRU cache.",
|
|
61
|
+
)
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
class EiaPetroleumStatusReportData(PetroleumStatusReportData):
|
|
65
|
+
"""EIA Petroleum Status Report Data Model."""
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
class EiaPetroleumStatusReportFetcher(
|
|
69
|
+
Fetcher[EiaPetroleumStatusReportQueryParams, list[EiaPetroleumStatusReportData]]
|
|
70
|
+
):
|
|
71
|
+
"""EIA Petroleum Status Report Fetcher."""
|
|
72
|
+
|
|
73
|
+
require_credentials = False
|
|
74
|
+
|
|
75
|
+
@staticmethod
|
|
76
|
+
def transform_query(params: dict[str, Any]) -> EiaPetroleumStatusReportQueryParams:
|
|
77
|
+
"""Transform the query parameters."""
|
|
78
|
+
# pylint: disable=import-outside-toplevel
|
|
79
|
+
from warnings import warn
|
|
80
|
+
|
|
81
|
+
category = params.get("category", "balance_sheet")
|
|
82
|
+
tables = WpsrTableMap.get(category, {})
|
|
83
|
+
_table = params.get("table", "")
|
|
84
|
+
|
|
85
|
+
if not _table:
|
|
86
|
+
_table = "stocks" if category == "weekly_estimates" else "all"
|
|
87
|
+
|
|
88
|
+
_tables = _table.split(",")
|
|
89
|
+
|
|
90
|
+
if len(_tables) == 1 and _tables[0] == "all" and category == "weekly_estimates":
|
|
91
|
+
raise OpenBBError(
|
|
92
|
+
ValueError(
|
|
93
|
+
f"'all' is not a supported choice for {category}. Please choose from: {list(tables)}"
|
|
94
|
+
)
|
|
95
|
+
)
|
|
96
|
+
|
|
97
|
+
if "all" in _tables and len(_tables) > 1:
|
|
98
|
+
_tables.remove("all")
|
|
99
|
+
warn("'all' cannot be used with other table choices. Ignoring 'all'.")
|
|
100
|
+
|
|
101
|
+
for table in _tables:
|
|
102
|
+
if table != "all" and table not in tables:
|
|
103
|
+
raise OpenBBError(
|
|
104
|
+
ValueError(
|
|
105
|
+
f"Invalid table choice: {table}. Valid choices for {category}: {list(tables)}"
|
|
106
|
+
)
|
|
107
|
+
)
|
|
108
|
+
|
|
109
|
+
params["table"] = ",".join(_tables)
|
|
110
|
+
|
|
111
|
+
return EiaPetroleumStatusReportQueryParams(**params)
|
|
112
|
+
|
|
113
|
+
@staticmethod
|
|
114
|
+
async def aextract_data(
|
|
115
|
+
query: EiaPetroleumStatusReportQueryParams,
|
|
116
|
+
credentials: Optional[dict[str, Any]],
|
|
117
|
+
**kwargs: Any,
|
|
118
|
+
) -> dict:
|
|
119
|
+
"""Extract the data from the EIA website."""
|
|
120
|
+
# pylint: disable=import-outside-toplevel
|
|
121
|
+
from openbb_us_eia.utils.helpers import download_excel_file
|
|
122
|
+
|
|
123
|
+
url = WpsrFileMap.get(query.category, "balance_sheet")
|
|
124
|
+
|
|
125
|
+
try:
|
|
126
|
+
results = await download_excel_file(url, query.use_cache)
|
|
127
|
+
except OpenBBError as e:
|
|
128
|
+
raise OpenBBError(f"Error extracting data -> {e}") from e
|
|
129
|
+
|
|
130
|
+
return {"file": results}
|
|
131
|
+
|
|
132
|
+
@staticmethod
|
|
133
|
+
def transform_data(
|
|
134
|
+
query: EiaPetroleumStatusReportQueryParams,
|
|
135
|
+
data: dict,
|
|
136
|
+
**kwargs: Any,
|
|
137
|
+
) -> list[EiaPetroleumStatusReportData]:
|
|
138
|
+
"""Transform the data."""
|
|
139
|
+
# pylint: disable=import-outside-toplevel
|
|
140
|
+
import concurrent.futures # noqa
|
|
141
|
+
import re
|
|
142
|
+
from functools import lru_cache
|
|
143
|
+
from numpy import nan
|
|
144
|
+
from pandas import Categorical, ExcelFile, concat, read_excel
|
|
145
|
+
from warnings import warn
|
|
146
|
+
|
|
147
|
+
category = query.category
|
|
148
|
+
|
|
149
|
+
_tables = (
|
|
150
|
+
query.table.split(",") # type: ignore
|
|
151
|
+
if query.table
|
|
152
|
+
else ["stocks"] if category == "weekly_estimates" else ["all"]
|
|
153
|
+
)
|
|
154
|
+
all_tables = list(WpsrTableMap[category])
|
|
155
|
+
tables = all_tables if "all" in _tables else _tables
|
|
156
|
+
|
|
157
|
+
file = data.get("file")
|
|
158
|
+
|
|
159
|
+
if not isinstance(file, ExcelFile):
|
|
160
|
+
raise OpenBBError(
|
|
161
|
+
TypeError(f"Expected an ExcelFile object, got {type(file)} instead.")
|
|
162
|
+
)
|
|
163
|
+
|
|
164
|
+
dfs: list = []
|
|
165
|
+
|
|
166
|
+
def replace_data_strings(text):
|
|
167
|
+
"""Replace the table strings with sortable numbers."""
|
|
168
|
+
pattern = r"Data (\d):"
|
|
169
|
+
|
|
170
|
+
def replacer(match):
|
|
171
|
+
"""Replace the matched string with a sortable number."""
|
|
172
|
+
return f"Data 0{match.group(1)}:"
|
|
173
|
+
|
|
174
|
+
return re.sub(pattern, replacer, text)
|
|
175
|
+
|
|
176
|
+
@lru_cache(maxsize=128)
|
|
177
|
+
def read_excel_file(file, category, table):
|
|
178
|
+
"""Read the ExcelFile for the sheet name and flatten the table."""
|
|
179
|
+
sheet_name = WpsrTableMap[category][table]
|
|
180
|
+
table_name = read_excel(file, sheet_name, header=None, nrows=1).iloc[0, 1]
|
|
181
|
+
table_name = replace_data_strings(table_name)
|
|
182
|
+
df = read_excel(file, sheet_name, header=[1, 2], nrows=3)
|
|
183
|
+
symbols = df.columns.get_level_values(0).tolist()
|
|
184
|
+
titles = [
|
|
185
|
+
d.replace(".1", "") for d in df.columns.get_level_values(1).tolist()
|
|
186
|
+
]
|
|
187
|
+
title_map = dict(zip(symbols, titles))
|
|
188
|
+
df = read_excel(file, sheet_name, header=None, skiprows=3)
|
|
189
|
+
df.columns = [d.replace("Sourcekey", "date") for d in symbols]
|
|
190
|
+
df = df.melt(
|
|
191
|
+
id_vars="date",
|
|
192
|
+
value_vars=[d for d in df.columns if d != "date"],
|
|
193
|
+
var_name="symbol",
|
|
194
|
+
).dropna()
|
|
195
|
+
df = df.reset_index(drop=True)
|
|
196
|
+
df.loc[:, "title"] = df.symbol.map(title_map)
|
|
197
|
+
df.loc[:, "unit"] = df.title.map(lambda x: x.split(" (")[-1].split(")")[0])
|
|
198
|
+
units = [f"({d})" for d in df.unit.unique().tolist()]
|
|
199
|
+
for unit in units:
|
|
200
|
+
df.title = df.title.str.replace(unit, "", regex=False).str.strip()
|
|
201
|
+
df.loc[:, "table"] = table_name
|
|
202
|
+
df["order"] = df.groupby("date").cumcount() + 1
|
|
203
|
+
df = df[["date", "table", "symbol", "order", "title", "value", "unit"]]
|
|
204
|
+
df.symbol = Categorical(df.symbol, categories=symbols, ordered=True)
|
|
205
|
+
df = df.sort_values(["date", "symbol"])
|
|
206
|
+
df.date = df.date.dt.date
|
|
207
|
+
|
|
208
|
+
if query.start_date:
|
|
209
|
+
df = df[df.date >= query.start_date]
|
|
210
|
+
|
|
211
|
+
if query.end_date:
|
|
212
|
+
df = df[df.date <= query.end_date]
|
|
213
|
+
|
|
214
|
+
df = df.reset_index(drop=True)
|
|
215
|
+
|
|
216
|
+
if len(df) > 0:
|
|
217
|
+
dfs.append(df)
|
|
218
|
+
else:
|
|
219
|
+
warn(f"No data for table: {table}")
|
|
220
|
+
|
|
221
|
+
try:
|
|
222
|
+
with concurrent.futures.ThreadPoolExecutor() as executor:
|
|
223
|
+
executor.map(
|
|
224
|
+
lambda table: read_excel_file(file, category, table), tables
|
|
225
|
+
)
|
|
226
|
+
|
|
227
|
+
results = concat(dfs)
|
|
228
|
+
|
|
229
|
+
if len(results) < 1:
|
|
230
|
+
raise EmptyDataError("The data is empty.")
|
|
231
|
+
|
|
232
|
+
results = results.sort_values(by=["date", "table", "order"]).replace(
|
|
233
|
+
{nan: None}
|
|
234
|
+
)
|
|
235
|
+
|
|
236
|
+
return [
|
|
237
|
+
EiaPetroleumStatusReportData.model_validate(d)
|
|
238
|
+
for d in results.to_dict(orient="records")
|
|
239
|
+
]
|
|
240
|
+
except Exception as e: # pylint: disable=broad-except
|
|
241
|
+
raise OpenBBError(f"Error transforming the data -> {e}") from e
|
|
@@ -0,0 +1,284 @@
|
|
|
1
|
+
"""EIA Short Term Energy Outlook Model."""
|
|
2
|
+
|
|
3
|
+
# pylint: disable=unused-argument,too-many-branches,too-many-statements,too-many-locals
|
|
4
|
+
|
|
5
|
+
from typing import Any, Literal, Optional
|
|
6
|
+
from warnings import warn
|
|
7
|
+
|
|
8
|
+
from openbb_core.app.model.abstract.error import OpenBBError
|
|
9
|
+
from openbb_core.provider.abstract.fetcher import Fetcher
|
|
10
|
+
from openbb_core.provider.standard_models.short_term_energy_outlook import (
|
|
11
|
+
ShortTermEnergyOutlookData,
|
|
12
|
+
ShortTermEnergyOutlookQueryParams,
|
|
13
|
+
)
|
|
14
|
+
from openbb_core.provider.utils.descriptions import QUERY_DESCRIPTIONS
|
|
15
|
+
from openbb_core.provider.utils.errors import EmptyDataError
|
|
16
|
+
from openbb_us_eia.utils.constants import (
|
|
17
|
+
SteoTableMap,
|
|
18
|
+
SteoTableNames,
|
|
19
|
+
SteoTableType,
|
|
20
|
+
)
|
|
21
|
+
from pydantic import Field
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
class EiaShortTermEnergyOutlookQueryParams(ShortTermEnergyOutlookQueryParams):
|
|
25
|
+
"""EIA Short Term Energy Outlook Query Parameters.
|
|
26
|
+
|
|
27
|
+
Monthly short term (18 month) projections using STEO model
|
|
28
|
+
|
|
29
|
+
Source: www.eia.gov/steo/
|
|
30
|
+
"""
|
|
31
|
+
|
|
32
|
+
__json_schema_extra__ = {
|
|
33
|
+
"symbol": {
|
|
34
|
+
"multiple_items_allowed": True,
|
|
35
|
+
},
|
|
36
|
+
"table": {
|
|
37
|
+
"multiple_items_allowed": False,
|
|
38
|
+
"choices": list(SteoTableNames),
|
|
39
|
+
},
|
|
40
|
+
"frequency": {
|
|
41
|
+
"multiple_items_allowed": False,
|
|
42
|
+
"choices": ["month", "quarter", "annual"],
|
|
43
|
+
},
|
|
44
|
+
}
|
|
45
|
+
symbol: Optional[str] = Field(
|
|
46
|
+
default=None,
|
|
47
|
+
description=QUERY_DESCRIPTIONS.get("symbol", "")
|
|
48
|
+
+ " If provided, overrides the 'table' parameter to return only the specified symbol from the STEO API.",
|
|
49
|
+
)
|
|
50
|
+
table: SteoTableType = Field(
|
|
51
|
+
default="01",
|
|
52
|
+
description="The specific table within the STEO dataset. Default is '01'."
|
|
53
|
+
+ " When 'symbol' is provided, this parameter is ignored."
|
|
54
|
+
+ "\n 01: US Energy Markets Summary"
|
|
55
|
+
+ "\n 02: Nominal Energy Prices"
|
|
56
|
+
+ "\n 03a: World Petroleum and Other Liquid Fuels Production, Consumption, and Inventories"
|
|
57
|
+
+ "\n 03b: Non-OPEC Petroleum and Other Liquid Fuels Production"
|
|
58
|
+
+ "\n 03c: World Petroleum and Other Liquid Fuels Production"
|
|
59
|
+
+ "\n 03d: World Crude Oil Production"
|
|
60
|
+
+ "\n 03e: World Petroleum and Other Liquid Fuels Consumption"
|
|
61
|
+
+ "\n 04a: US Petroleum and Other Liquid Fuels Supply, Consumption, and Inventories"
|
|
62
|
+
+ "\n 04b: US Hydrocarbon Gas Liquids (HGL) and Petroleum Refinery Balances"
|
|
63
|
+
+ "\n 04c: US Regional Motor Gasoline Prices and Inventories"
|
|
64
|
+
+ "\n 04d: US Biofuel Supply, Consumption, and Inventories"
|
|
65
|
+
+ "\n 05a: US Natural Gas Supply, Consumption, and Inventories"
|
|
66
|
+
+ "\n 05b: US Regional Natural Gas Prices"
|
|
67
|
+
+ "\n 06: US Coal Supply, Consumption, and Inventories"
|
|
68
|
+
+ "\n 07a: US Electricity Industry Overview"
|
|
69
|
+
+ "\n 07b: US Regional Electricity Retail Sales"
|
|
70
|
+
+ "\n 07c: US Regional Electricity Prices"
|
|
71
|
+
+ "\n 07d1: US Regional Electricity Generation, Electric Power Sector"
|
|
72
|
+
+ "\n 07d2: US Regional Electricity Generation, Electric Power Sector, continued"
|
|
73
|
+
+ "\n 07e: US Electricity Generating Capacity"
|
|
74
|
+
+ "\n 08: US Renewable Energy Consumption"
|
|
75
|
+
+ "\n 09a: US Macroeconomic Indicators and CO2 Emissions"
|
|
76
|
+
+ "\n 09b: US Regional Macroeconomic Data"
|
|
77
|
+
+ "\n 09c: US Regional Weather Data"
|
|
78
|
+
+ "\n 10a: Drilling Productivity Metrics"
|
|
79
|
+
+ "\n 10b: Crude Oil and Natural Gas Production from Shale and Tight Formations",
|
|
80
|
+
)
|
|
81
|
+
frequency: Literal["month", "quarter", "annual"] = Field(
|
|
82
|
+
default="month",
|
|
83
|
+
description="The frequency of the data. Default is 'month'.",
|
|
84
|
+
)
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
class EiaShortTermEnergyOutlookData(ShortTermEnergyOutlookData):
|
|
88
|
+
"""EIA Short Term Energy Outlook Data Model."""
|
|
89
|
+
|
|
90
|
+
__alias_dict__ = {
|
|
91
|
+
"date": "period",
|
|
92
|
+
"symbol": "seriesId",
|
|
93
|
+
"title": "seriesDescription",
|
|
94
|
+
}
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
class EiaShortTermEnergyOutlookFetcher(
|
|
98
|
+
Fetcher[EiaShortTermEnergyOutlookQueryParams, list[EiaShortTermEnergyOutlookData]]
|
|
99
|
+
):
|
|
100
|
+
"""EIA Short Term Energy Outlook Fetcher."""
|
|
101
|
+
|
|
102
|
+
@staticmethod
|
|
103
|
+
def transform_query(params: dict[str, Any]) -> EiaShortTermEnergyOutlookQueryParams:
|
|
104
|
+
"""Transform the query parameters."""
|
|
105
|
+
return EiaShortTermEnergyOutlookQueryParams(**params)
|
|
106
|
+
|
|
107
|
+
@staticmethod
|
|
108
|
+
async def aextract_data(
|
|
109
|
+
query: EiaShortTermEnergyOutlookQueryParams,
|
|
110
|
+
credentials: Optional[dict[str, Any]],
|
|
111
|
+
**kwargs: Any,
|
|
112
|
+
) -> list[dict]:
|
|
113
|
+
"""Extract the data from the EIA API."""
|
|
114
|
+
# pylint: disable=import-outside-toplevel
|
|
115
|
+
import asyncio # noqa
|
|
116
|
+
from openbb_core.provider.utils.helpers import amake_request
|
|
117
|
+
from openbb_us_eia.utils.helpers import response_callback
|
|
118
|
+
|
|
119
|
+
api_key = credentials.get("eia_api_key") if credentials else ""
|
|
120
|
+
frequency_dict = {
|
|
121
|
+
"month": "monthly",
|
|
122
|
+
"quarter": "quarterly",
|
|
123
|
+
"annual": "annual",
|
|
124
|
+
}
|
|
125
|
+
frequency = frequency_dict[query.frequency]
|
|
126
|
+
base_url = f"https://api.eia.gov/v2/steo/data/?api_key={api_key}&frequency={frequency}&data[0]=value"
|
|
127
|
+
urls: list[str] = []
|
|
128
|
+
start_date: str = ""
|
|
129
|
+
end_date: str = ""
|
|
130
|
+
|
|
131
|
+
# Format the dates based on the frequency.
|
|
132
|
+
def resample_to_quarter(dt) -> str:
|
|
133
|
+
"""Resample a date to a string formatted as 'YYYY-QX'."""
|
|
134
|
+
year = dt.year
|
|
135
|
+
quarter = (dt.month - 1) // 3 + 1
|
|
136
|
+
return f"{year}-Q{quarter}"
|
|
137
|
+
|
|
138
|
+
if query.start_date is not None and frequency == "monthly":
|
|
139
|
+
start_date = f"&start={query.start_date.strftime('%Y-%m')}"
|
|
140
|
+
elif query.start_date is not None and frequency == "quarterly":
|
|
141
|
+
start_date = f"&start={resample_to_quarter(query.start_date)}"
|
|
142
|
+
elif query.start_date is not None and frequency == "annual":
|
|
143
|
+
start_date = f"&start={query.start_date.strftime('%Y')}"
|
|
144
|
+
|
|
145
|
+
if query.end_date is not None and frequency == "monthly":
|
|
146
|
+
end_date = f"&end={query.end_date.strftime('%Y-%m')}"
|
|
147
|
+
elif query.end_date is not None and frequency == "quarterly":
|
|
148
|
+
end_date = f"&end={resample_to_quarter(query.end_date)}"
|
|
149
|
+
elif query.end_date is not None and frequency == "annual":
|
|
150
|
+
end_date = f"&end={query.end_date.strftime('%Y')}"
|
|
151
|
+
|
|
152
|
+
# We chunk the request to avoid pagination and make the query execution faster.
|
|
153
|
+
symbols = (
|
|
154
|
+
query.symbol.upper().split(",")
|
|
155
|
+
if query.symbol
|
|
156
|
+
else [d.upper() for d in SteoTableMap[query.table]]
|
|
157
|
+
)
|
|
158
|
+
seen = set()
|
|
159
|
+
unique_symbols: list = []
|
|
160
|
+
for symbol in symbols:
|
|
161
|
+
if symbol not in seen:
|
|
162
|
+
unique_symbols.append(symbol)
|
|
163
|
+
seen.add(symbol)
|
|
164
|
+
symbols = unique_symbols
|
|
165
|
+
|
|
166
|
+
def encode_symbols(symbol: str):
|
|
167
|
+
"""Encode a chunk of symbols to be used in a URL"""
|
|
168
|
+
prefix = "&facets[seriesId][]="
|
|
169
|
+
return prefix + symbol.upper()
|
|
170
|
+
|
|
171
|
+
for i in range(0, len(symbols), 10):
|
|
172
|
+
url_symbols: str = ""
|
|
173
|
+
symbols_chunk = symbols[i : i + 10]
|
|
174
|
+
for symbol in symbols_chunk:
|
|
175
|
+
url_symbols += encode_symbols(symbol)
|
|
176
|
+
url = f"{base_url}{url_symbols}{start_date}{end_date}&offset=0&length=5000"
|
|
177
|
+
urls.append(url)
|
|
178
|
+
|
|
179
|
+
results: list[dict] = []
|
|
180
|
+
messages: list[str] = []
|
|
181
|
+
|
|
182
|
+
async def get_one(url):
|
|
183
|
+
"""Response callback function."""
|
|
184
|
+
res = await amake_request(url, response_callback=response_callback)
|
|
185
|
+
data = res.get("response", {}).get("data", []) # type: ignore
|
|
186
|
+
if not data:
|
|
187
|
+
series_id = (
|
|
188
|
+
res.get("request", {}) # type: ignore
|
|
189
|
+
.get("params", {})
|
|
190
|
+
.get("facets", {})
|
|
191
|
+
.get("seriesId", [])
|
|
192
|
+
)
|
|
193
|
+
masked_url = url.replace(api_key, "API_KEY")
|
|
194
|
+
messages.append(f"No data returned for {series_id or masked_url}")
|
|
195
|
+
if data:
|
|
196
|
+
results.extend(data)
|
|
197
|
+
response_total = int(res.get("response", {}).get("total", 0)) # type: ignore
|
|
198
|
+
n_results = len(data)
|
|
199
|
+
# After conservatively chunking the request, we may still need to paginate.
|
|
200
|
+
# This is mostly out of an abundance of caution.
|
|
201
|
+
if response_total > 5000 and n_results == 5000:
|
|
202
|
+
offset = 5000
|
|
203
|
+
url = url.replace("&offset=0", f"&offset={offset}")
|
|
204
|
+
while n_results < response_total:
|
|
205
|
+
additional_response = await amake_request(url)
|
|
206
|
+
additional_data = additional_response.get("response", {}).get( # type: ignore
|
|
207
|
+
"data", []
|
|
208
|
+
)
|
|
209
|
+
if not additional_data:
|
|
210
|
+
series_id = (
|
|
211
|
+
res.get("request", {}) # type: ignore
|
|
212
|
+
.get("params", {})
|
|
213
|
+
.get("facets", {})
|
|
214
|
+
.get("seriesId", [])
|
|
215
|
+
)
|
|
216
|
+
masked_url = url.replace(api_key, "API_KEY")
|
|
217
|
+
messages.append(
|
|
218
|
+
f"No additional data returned for {series_id or masked_url}"
|
|
219
|
+
)
|
|
220
|
+
if additional_data:
|
|
221
|
+
results.extend(additional_data)
|
|
222
|
+
n_results += len(additional_data)
|
|
223
|
+
url = url.replace(f"&offset={offset}", f"&offset={offset+5000}")
|
|
224
|
+
offset += 5000
|
|
225
|
+
|
|
226
|
+
try:
|
|
227
|
+
await asyncio.gather(*[get_one(url) for url in urls])
|
|
228
|
+
except Exception as e:
|
|
229
|
+
raise OpenBBError(f"Error fetching data from the EIA API -> {e}") from e
|
|
230
|
+
|
|
231
|
+
if not results and not messages:
|
|
232
|
+
raise EmptyDataError(
|
|
233
|
+
"The request was returned empty with no error messages."
|
|
234
|
+
)
|
|
235
|
+
if not results and messages:
|
|
236
|
+
raise OpenBBError("\n".join(messages))
|
|
237
|
+
if results and messages:
|
|
238
|
+
warn("\n".join(messages))
|
|
239
|
+
|
|
240
|
+
return results
|
|
241
|
+
|
|
242
|
+
@staticmethod
|
|
243
|
+
def transform_data(
|
|
244
|
+
query: EiaShortTermEnergyOutlookQueryParams,
|
|
245
|
+
data: list[dict],
|
|
246
|
+
**kwargs: Any,
|
|
247
|
+
) -> list[EiaShortTermEnergyOutlookData]:
|
|
248
|
+
"""Transform the data."""
|
|
249
|
+
# pylint: disable=import-outside-toplevel
|
|
250
|
+
from pandas import Categorical, DataFrame, to_datetime
|
|
251
|
+
|
|
252
|
+
symbols = (
|
|
253
|
+
query.symbol.upper().split(",")
|
|
254
|
+
if query.symbol
|
|
255
|
+
else [d.upper() for d in SteoTableMap[query.table]]
|
|
256
|
+
)
|
|
257
|
+
seen = set()
|
|
258
|
+
unique_symbols: list = []
|
|
259
|
+
for symbol in symbols:
|
|
260
|
+
if symbol not in seen:
|
|
261
|
+
unique_symbols.append(symbol)
|
|
262
|
+
seen.add(symbol)
|
|
263
|
+
symbols = unique_symbols
|
|
264
|
+
|
|
265
|
+
table = query.table
|
|
266
|
+
df = DataFrame(data)
|
|
267
|
+
df.period = to_datetime(df.period).dt.date
|
|
268
|
+
df.seriesId = Categorical(df.seriesId, categories=symbols, ordered=True)
|
|
269
|
+
df = df.sort_values(["period", "seriesId"])
|
|
270
|
+
df = df.reset_index(drop=True)
|
|
271
|
+
returned_symbols = df.seriesId.unique().tolist()
|
|
272
|
+
missing_symbols = [s for s in symbols if s not in returned_symbols]
|
|
273
|
+
|
|
274
|
+
if query.symbol and missing_symbols:
|
|
275
|
+
warn(f"No data was returned for: {', '.join(missing_symbols)}")
|
|
276
|
+
|
|
277
|
+
if not query.symbol:
|
|
278
|
+
df["order"] = df.groupby("period").cumcount() + 1
|
|
279
|
+
df["table"] = (
|
|
280
|
+
f"STEO - {table.replace('0', '') if table[0] == '0' else table}: {SteoTableNames[table]}"
|
|
281
|
+
)
|
|
282
|
+
records = df.to_dict(orient="records")
|
|
283
|
+
|
|
284
|
+
return [EiaShortTermEnergyOutlookData.model_validate(d) for d in records]
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
"""OpenBB EIA Provider Utilities."""
|