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.
@@ -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."""