fix: unify mypy environments for local checks and CI (#1291)

The isolated pre-commit mypy hook previously omitted runtime type information that
make mypy used, hiding errors involving dependencies such as Pydantic and Pendulum.
Makefile, pre-commit and CI now run the same full-project typing policy in the
development environment defined by uv.lock.

- Use uv run --locked --exact --extra dev and the same mypy arguments for Makefile
  and the local hook. Check all of src and tests, including on configuration-only
  changes.
- Pin Python 3.13 for local development and the pre-commit CI job, and install the
  locked pre-commit version in CI.
- Disable incremental analysis because existing Pendulum cache state changes mypy 2.3.1
  diagnostics. Document the policy, the performance tradeoff and the existing typing debt.
- Add a regression test that exercises Makefile, the hook and the CI command in a
  temporary project, accepting valid dependency types and detecting deliberate
  Pydantic/Pendulum assignment errors.

Resolve the newly detected mypy diagnostics.

- Enable the numpydantic and Pydantic mypy plugins, retaining strict Pydantic
  constructor typing with init_typed = true. Validate raw/coercible payloads through model_validate.
- Propagate concrete record, provider and time-window types through generic collections,
  factories and lookup methods. Preserve runtime field inspection and generated time-window
  documentation.
- Align Pendulum annotations with actual factory/arithmetic results while retaining Pydantic
  validation adapters at runtime. Correct optional values, array boundaries, REST handlers
  and plotting interfaces.
- Add pinned scipy-stubs and types-psutil, update uv.lock, and supply the plugins' dependencies.
- Add runtime regression coverage for validated path defaults, normalized time-series metadata,
  generic field inspection, invalid timestamps and unsupported provider imports.

Runtime and compatibility details:

- Validate path defaults as Path objects while retaining raw string defaults needed by
  migration serialization with exclude_defaults.
- Normalize feed-in tariff lists and default charge rates to NumPy arrays; reject missing
  timestamps/uninitialized values explicitly. Importing into a provider without import support
  returns HTTP 400.
- Public JSON schemas and OpenAPI structure match main (excluding the generated version).

Signed-off-by: dr-dimitry

Signed-off-by: dr-dimitry
Signed-off-by: Bobby Noelte <b0661n0e17e@gmail.com>
Co-authored-by: dr-dimitri <87113560+dr-dimitri@users.noreply.github.com>
Co-authored-by: Normann <github@koldrack.com>
This commit is contained in:
Bobby Noelte
2026-09-10 23:20:35 +02:00
committed by GitHub
co-authored by dr-dimitri Normann
parent 5b584cbb57
commit 1abdd345c4
114 changed files with 2217 additions and 1322 deletions
+1 -1
View File
@@ -69,7 +69,7 @@ def adapter_providers() -> list[Union["HomeAssistantAdapter", "NodeREDAdapter"]]
]
class Adapter(AdapterContainer):
class Adapter(AdapterContainer[HomeAssistantAdapter | NodeREDAdapter]):
"""Adapter container to manage multiple adapter providers."""
providers: list[
+8 -5
View File
@@ -2,7 +2,7 @@
import asyncio
from abc import abstractmethod
from typing import Any, Optional
from typing import Any, Generic, Optional, TypeVar
from loguru import logger
from pydantic import (
@@ -102,18 +102,21 @@ class AdapterProvider(SingletonMixin, ConfigMixin, MeasurementMixin, StartMixin,
await self._update_data()
class AdapterContainer(SingletonMixin, ConfigMixin, PydanticBaseModel):
AdapterProviderT = TypeVar("AdapterProviderT", bound=AdapterProvider)
class AdapterContainer(SingletonMixin, ConfigMixin, PydanticBaseModel, Generic[AdapterProviderT]):
"""A container for managing multiple adapter provider instances.
This class enables to control multiple adapter providers
"""
providers: list[AdapterProvider] = Field(
providers: list[AdapterProviderT] = Field(
default_factory=list, json_schema_extra={"description": "List of adapter providers"}
)
@field_validator("providers")
def check_providers(cls, value: list[AdapterProvider]) -> list[AdapterProvider]:
def check_providers(cls, value: list[AdapterProviderT]) -> list[AdapterProviderT]:
# Check each item in the list
for item in value:
if not isinstance(item, AdapterProvider):
@@ -149,7 +152,7 @@ class AdapterContainer(SingletonMixin, ConfigMixin, PydanticBaseModel):
return
super().__init__(*args, **kwargs)
def provider_by_id(self, provider_id: str) -> AdapterProvider:
def provider_by_id(self, provider_id: str) -> AdapterProviderT:
"""Retrieves an adapter provider by its unique identifier.
This method searches through the list of all available providers and
+21 -5
View File
@@ -145,7 +145,13 @@ class HomeAssistantAdapterCommonSettings(SettingsBaseModel):
"""Entity IDs available at Home Assistant."""
try:
adapter_eos = get_adapter()
result = adapter_eos.provider_by_id("HomeAssistant").get_homeassistant_entity_ids()
provider = adapter_eos.provider_by_id("HomeAssistant")
except Exception:
return []
if not isinstance(provider, HomeAssistantAdapter):
raise TypeError("HomeAssistant provider must be a HomeAssistantAdapter")
try:
result = provider.get_homeassistant_entity_ids()
except Exception:
return []
return result
@@ -156,7 +162,13 @@ class HomeAssistantAdapterCommonSettings(SettingsBaseModel):
"""Entity IDs for optimization solution available at EOS."""
try:
adapter_eos = get_adapter()
result = adapter_eos.provider_by_id("HomeAssistant").get_eos_solution_entity_ids()
provider = adapter_eos.provider_by_id("HomeAssistant")
except Exception:
return []
if not isinstance(provider, HomeAssistantAdapter):
raise TypeError("HomeAssistant provider must be a HomeAssistantAdapter")
try:
result = provider.get_eos_solution_entity_ids()
except Exception:
return []
return result
@@ -167,9 +179,13 @@ class HomeAssistantAdapterCommonSettings(SettingsBaseModel):
"""Entity IDs for energy management instructions available at EOS."""
try:
adapter_eos = get_adapter()
result = adapter_eos.provider_by_id(
"HomeAssistant"
).get_eos_device_instruction_entity_ids()
provider = adapter_eos.provider_by_id("HomeAssistant")
except Exception:
return []
if not isinstance(provider, HomeAssistantAdapter):
raise TypeError("HomeAssistant provider must be a HomeAssistantAdapter")
try:
result = provider.get_eos_device_instruction_entity_ids()
except Exception:
return []
return result
+20 -5
View File
@@ -14,7 +14,7 @@ import os
import sys
import tempfile
from pathlib import Path
from typing import Any, ClassVar, Optional, Type, Union
from typing import Any, Callable, ClassVar, Optional, Type, Union
import pydantic_settings
from loguru import logger
@@ -122,6 +122,10 @@ def default_data_folder_path() -> Path:
class GeneralSettings(SettingsBaseModel):
"""General settings."""
# Legacy configuration-path metadata populated by ConfigEOS._setup_config_file.
_config_file_path: ClassVar[Path | None] = None
_config_folder_path: ClassVar[Path | None] = None
config_save_mode: ConfigSaveMode = Field(
default=ConfigSaveMode.AUTOMATIC,
json_schema_extra={
@@ -161,8 +165,11 @@ class GeneralSettings(SettingsBaseModel):
},
)
# Validate this raw default to Path. Retain the string so
# exclude_defaults preserves the output path in migrated configurations.
data_output_subpath: Optional[Path] = Field(
default="output",
validate_default=True,
json_schema_extra={"description": "Sub-path for the EOS output data folder."},
)
@@ -402,14 +409,15 @@ class ConfigEOS(SingletonMixin, SettingsEOSDefaults):
return True
@classmethod
def settings_customise_sources(
# Pydantic Settings accepts zero-argument callables as well as source objects.
def settings_customise_sources( # type: ignore[override]
cls,
settings_cls: Type[pydantic_settings.BaseSettings],
init_settings: pydantic_settings.PydanticBaseSettingsSource,
env_settings: pydantic_settings.PydanticBaseSettingsSource,
dotenv_settings: pydantic_settings.PydanticBaseSettingsSource,
file_secret_settings: pydantic_settings.PydanticBaseSettingsSource,
) -> tuple[pydantic_settings.PydanticBaseSettingsSource, ...]:
) -> tuple[pydantic_settings.PydanticBaseSettingsSource | Callable[[], dict[str, Any]], ...]:
"""Customizes the order and handling of settings sources for a pydantic_settings.BaseSettings subclass.
This method determines the sources for application configuration settings, including
@@ -790,7 +798,11 @@ class ConfigEOS(SingletonMixin, SettingsEOSDefaults):
required by ``self._setup()``.
OSError: If reading the backup file fails due to I/O issues.
"""
backup_file_path = self.general.config_file_path.with_suffix(f".{backup_id}")
config_file_path = self.general.config_file_path
# Configuration setup initializes this path; should never raise.
if config_file_path is None:
raise AssertionError("Configuration file path is not initialized")
backup_file_path = config_file_path.with_suffix(f".{backup_id}")
if not backup_file_path.exists():
error_msg = f"Configuration backup `{backup_id}` not found."
logger.error(error_msg)
@@ -823,7 +835,10 @@ class ConfigEOS(SingletonMixin, SettingsEOSDefaults):
"""
result: dict[str, dict[str, Any]] = {}
base_path: Path = self.general.config_file_path
base_path = self.general.config_file_path
# Configuration setup initializes this path; should never raise.
if base_path is None:
raise AssertionError("Configuration file path is not initialized")
parent = base_path.parent
stem = base_path.stem
+18 -12
View File
@@ -4,7 +4,7 @@ import calendar
import os
import sys
from enum import StrEnum
from typing import Any, ClassVar, Iterator, Optional, Union
from typing import Any, ClassVar, Generic, Iterator, Optional, TypeVar, Union
import numpy as np
import pandas as pd
@@ -409,19 +409,23 @@ class TimeWindow(SettingsBaseModel):
return self.duration
class TimeWindowSequence(SettingsBaseModel):
TimeWindowT = TypeVar("TimeWindowT", bound=TimeWindow)
class TimeWindowSequence(SettingsBaseModel, Generic[TimeWindowT]):
"""Model representing a sequence of time windows with collective operations.
Manages multiple TimeWindow objects and provides methods to work with them
as a cohesive unit for scheduling and availability checking.
"""
windows: list[TimeWindow] = Field(
windows: list[TimeWindowT] = Field(
default_factory=list,
json_schema_extra={"description": "List of TimeWindow objects that make up this sequence."},
)
def __iter__(self) -> Iterator[TimeWindow]:
# EOS collections iterate over their elements instead of BaseModel field/value pairs.
def __iter__(self) -> Iterator[TimeWindowT]: # type: ignore[override]
"""Allow iteration over the time windows."""
return iter(self.windows)
@@ -429,7 +433,7 @@ class TimeWindowSequence(SettingsBaseModel):
"""Return the number of time windows in the sequence."""
return len(self.windows)
def __getitem__(self, index: int) -> TimeWindow:
def __getitem__(self, index: int) -> TimeWindowT:
"""Allow indexing into the time windows."""
return self.windows[index]
@@ -536,7 +540,9 @@ class TimeWindowSequence(SettingsBaseModel):
total += d
return total
def get_applicable_windows(self, reference_date: Optional[DateTime] = None) -> list[TimeWindow]:
def get_applicable_windows(
self, reference_date: Optional[DateTime] = None
) -> list[TimeWindowT]:
"""Get all windows that apply to the given reference date.
Args:
@@ -556,7 +562,7 @@ class TimeWindowSequence(SettingsBaseModel):
def find_windows_for_duration(
self, duration: Duration, reference_date: Optional[DateTime] = None
) -> list[TimeWindow]:
) -> list[TimeWindowT]:
"""Find all windows that can accommodate the given duration.
Args:
@@ -575,7 +581,7 @@ class TimeWindowSequence(SettingsBaseModel):
def get_all_possible_start_times(
self, duration: Duration, reference_date: Optional[DateTime] = None
) -> list[tuple[DateTime, DateTime, TimeWindow]]:
) -> list[tuple[DateTime, DateTime, TimeWindowT]]:
"""Get all possible start time ranges for a duration across all windows.
Args:
@@ -739,7 +745,7 @@ class TimeWindowSequence(SettingsBaseModel):
dtype=np.float64,
)
def add_window(self, window: TimeWindow) -> None:
def add_window(self, window: TimeWindowT) -> None:
"""Add a new time window to the sequence.
Args:
@@ -747,7 +753,7 @@ class TimeWindowSequence(SettingsBaseModel):
"""
self.windows.append(window)
def remove_window(self, index: int) -> TimeWindow:
def remove_window(self, index: int) -> TimeWindowT:
"""Remove a time window from the sequence by index.
Args:
@@ -781,7 +787,7 @@ class TimeWindowSequence(SettingsBaseModel):
if reference_date is None:
reference_date = pendulum.today()
def sort_key(window: TimeWindow) -> tuple[int, DateTime]:
def sort_key(window: TimeWindowT) -> tuple[int, DateTime]:
start_time = window.earliest_start_time(Duration(), reference_date)
if start_time is None:
return (1, reference_date)
@@ -806,7 +812,7 @@ class ValueTimeWindow(TimeWindow):
)
class ValueTimeWindowSequence(TimeWindowSequence):
class ValueTimeWindowSequence(TimeWindowSequence[ValueTimeWindow]):
"""Sequence of value time windows.
This model specializes `TimeWindowSequence` to ensure that all
+24 -17
View File
@@ -25,6 +25,7 @@ from typing import (
Optional,
ParamSpec,
TypeVar,
cast,
)
import cachebox
@@ -46,7 +47,8 @@ from akkudoktoreos.utils.datetimeutil import (
# ---------------------------------
# Define a type variable for methods and functions
TCallable = TypeVar("TCallable", bound=Callable[..., Any])
Param = ParamSpec("Param")
RetType = TypeVar("RetType")
def cache_energy_management_store_callback(event: int, key: Any, value: Any) -> None:
@@ -195,7 +197,9 @@ class CacheEnergyManagementStore(SingletonMixin):
raise AttributeError(f"'{self.cache.__class__.__name__}' object has no method 'clear'")
def cache_energy_management(callable: TCallable) -> TCallable:
def cache_energy_management(
func: Callable[Param, RetType],
) -> Callable[Param, RetType]:
"""Decorator for in memory caching the result of a callable.
This decorator caches the method or function's result in `CacheEnergyManagementStore`,
@@ -203,7 +207,7 @@ def cache_energy_management(callable: TCallable) -> TCallable:
next energy management start.
Args:
callable (Callable): The function or method to be decorated.
func (Callable): The function or method to be decorated.
Returns:
Callable: The wrapped function with caching functionality.
@@ -218,24 +222,22 @@ def cache_energy_management(callable: TCallable) -> TCallable:
"""
@cachebox.cached(
cache=CacheEnergyManagementStore().cache, callback=cache_energy_management_store_callback
)
@functools.wraps(callable)
def wrapper(*args: Any, **kwargs: Any) -> Any:
result = callable(*args, **kwargs)
return result
@functools.wraps(func)
def wrapper(*args: Param.args, **kwargs: Param.kwargs) -> RetType:
return func(*args, **kwargs)
return wrapper
cached_wrapper = cachebox.cached(
cache=CacheEnergyManagementStore().cache,
callback=cache_energy_management_store_callback,
)(wrapper)
return cast(Callable[Param, RetType], cached_wrapper)
# ---------------------------------
# Cache File Management
# ---------------------------------
Param = ParamSpec("Param")
RetType = TypeVar("RetType")
def cache_clear(clear_all: Optional[bool] = None) -> None:
"""Cleanup expired cache files."""
@@ -742,6 +744,9 @@ class CacheFileStore(ConfigMixin, SingletonMixin):
if clear_all:
clear_file = True
else:
# Initialized above when clear_all is false; should never raise.
if before_datetime is None:
raise AssertionError("Cache expiry threshold is not initialized")
clear_file = compare_datetimes(cache_item.until_datetime, before_datetime).lt
if clear_file:
@@ -782,9 +787,11 @@ class CacheFileStore(ConfigMixin, SingletonMixin):
with self._store_lock:
store_current = {}
for key, record in self._store.items():
ttl_duration = record.ttl_duration
if ttl_duration:
ttl_duration = ttl_duration.total_seconds()
ttl_duration = (
record.ttl_duration.total_seconds()
if record.ttl_duration
else record.ttl_duration
)
store_current[key] = {
# Convert file-like objects to file paths for serialization
"cache_file": self._get_file_path(record.cache_file),
+2
View File
@@ -14,8 +14,10 @@ from akkudoktoreos.config.configabc import SettingsBaseModel
class CacheCommonSettings(SettingsBaseModel):
"""Cache Configuration."""
# Retain the raw serialized default for exclude_defaults compatibility.
subpath: Optional[Path] = Field(
default="cache",
validate_default=True,
json_schema_extra={"description": "Sub-path for the EOS cache data directory."},
)
+54 -31
View File
@@ -20,11 +20,14 @@ from typing import (
TYPE_CHECKING,
Any,
Dict,
Generic,
Iterator,
Optional,
Tuple,
Type,
TypeVar,
Union,
cast,
get_args,
overload,
)
@@ -289,7 +292,8 @@ class DataRecord(DataABC, MutableMapping):
except AttributeError:
raise KeyError(f"'{key}' is not a recognized field.")
def __iter__(self) -> Iterator[str]:
# EOS collections iterate over their elements instead of BaseModel field/value pairs.
def __iter__(self) -> Iterator[str]: # type: ignore[override]
"""Iterate over the field names in the data record.
Returns:
@@ -441,7 +445,10 @@ class DataRecord(DataABC, MutableMapping):
# ==================== DataSequence ====================
class DataSequence(DataABC, DatabaseRecordProtocolMixin[DataRecord]):
DataRecordT = TypeVar("DataRecordT", bound=DataRecord)
class DataSequence(DataABC, DatabaseRecordProtocolMixin[DataRecordT], Generic[DataRecordT]):
"""A managed sequence of DataRecord instances with time series behavior.
The DataSequence class provides an ordered, mutable collection of DataRecord
@@ -488,7 +495,7 @@ class DataSequence(DataABC, DatabaseRecordProtocolMixin[DataRecord]):
"""
# To be overloaded by derived classes.
records: list[DataRecord] = Field(
records: list[DataRecordT] = Field(
default_factory=list, json_schema_extra={"description": "List of data records"}
)
@@ -553,7 +560,7 @@ class DataSequence(DataABC, DatabaseRecordProtocolMixin[DataRecord]):
f"Key '{key}' is not in writable record keys: {self.record_keys_writable}"
)
def _validate_record(self, value: DataRecord) -> None:
def _validate_record(self, value: DataRecordT) -> None:
"""Check if the provided value is a valid DataRecord with compatible keys.
Args:
@@ -629,7 +636,7 @@ class DataSequence(DataABC, DatabaseRecordProtocolMixin[DataRecord]):
return self.record_class().record_keys_writable()
@classmethod
def record_class(cls) -> Type:
def record_class(cls) -> Type[DataRecordT]:
"""Get the class of the data record handled by this data sequence.
This method determines the class of the data record type associated with
@@ -646,6 +653,8 @@ class DataSequence(DataABC, DatabaseRecordProtocolMixin[DataRecord]):
field_info = cls.model_fields["records"]
# Get the list element type from the 'type_' attribute
list_element_type = get_args(field_info.annotation)[0]
if isinstance(list_element_type, TypeVar):
list_element_type = list_element_type.__bound__
if not isinstance(list_element_type(), DataRecord):
raise ValueError(
f"Data record must be an instance of DataRecord: '{list_element_type}'."
@@ -736,17 +745,18 @@ class DataSequence(DataABC, DatabaseRecordProtocolMixin[DataRecord]):
# Sequence methods
def __iter__(self) -> Iterator[DataRecord]:
# EOS collections iterate over their elements instead of BaseModel field/value pairs.
def __iter__(self) -> Iterator[DataRecordT]: # type: ignore[override]
"""Create an iterator for accessing DataRecords sequentially (memory only).
Returns:
Iterator[DataRecord]: An iterator for the records.
Iterator[DataRecordT]: An iterator for the records.
"""
return iter(self.records)
async def get_by_datetime(
self, target_datetime: DateTime, *, time_window: Optional[Duration] = None
) -> Optional[DataRecord]:
) -> Optional[DataRecordT]:
"""Get the record at the specified datetime, with an optional fallback search window.
Args:
@@ -770,7 +780,7 @@ class DataSequence(DataABC, DatabaseRecordProtocolMixin[DataRecord]):
async def get_nearest_by_datetime(
self, target_datetime: DateTime, time_window: Optional[Duration] = None
) -> Optional[DataRecord]:
) -> Optional[DataRecordT]:
"""Get the record nearest to the specified datetime within an optional time window.
Args:
@@ -800,7 +810,7 @@ class DataSequence(DataABC, DatabaseRecordProtocolMixin[DataRecord]):
# sync rw write access to data sequence, needs locking in case of use in async.
async def _insert_by_datetime(self, record: DataRecord) -> None:
async def _insert_by_datetime(self, record: DataRecordT) -> None:
"""Insert or merge a DataRecord into the sequence based on its datetime.
Internal implementation of `insert_by_datetime`. Callers must
@@ -822,8 +832,10 @@ class DataSequence(DataABC, DatabaseRecordProtocolMixin[DataRecord]):
"""
self._validate_record(record)
# Ensure datetime objects are normalized
record_date_time_timestamp = DatabaseTimestamp.from_datetime(record.date_time)
# _validate_record normalizes the timestamp, including a missing value.
record_date_time_timestamp = DatabaseTimestamp.from_datetime(
self._db_require_date_time(record)
)
avail_record = await self.db_get_record(record_date_time_timestamp)
if avail_record:
@@ -914,7 +926,9 @@ class DataSequence(DataABC, DatabaseRecordProtocolMixin[DataRecord]):
avail_record = await self.db_get_record(db_target)
if avail_record is None:
# Create a new DataRecord if none exists
new_record = self.record_class()(date_time=date_time, **{key: values[i]})
new_record = self.record_class().model_validate(
{"date_time": date_time, key: values[i]}
)
await self.db_insert_record(new_record)
else:
# Update existing record's specified key
@@ -949,7 +963,9 @@ class DataSequence(DataABC, DatabaseRecordProtocolMixin[DataRecord]):
avail_record = await self.db_get_record(db_target)
if avail_record is None:
# Create a new DataRecord if none exists
new_record = self.record_class()(date_time=date_time, **{key: value})
new_record = self.record_class().model_validate(
{"date_time": date_time, key: value}
)
await self.db_insert_record(new_record)
else:
# Update existing record's specified key
@@ -958,7 +974,7 @@ class DataSequence(DataABC, DatabaseRecordProtocolMixin[DataRecord]):
# data sequence access usable also for async access
async def insert_by_datetime(self, record: DataRecord) -> None:
async def insert_by_datetime(self, record: DataRecordT) -> None:
"""Insert or merge a DataRecord into the sequence based on its date.
If a record with the same date exists, merges new data fields with the existing record.
@@ -1010,7 +1026,7 @@ class DataSequence(DataABC, DatabaseRecordProtocolMixin[DataRecord]):
start_datetime: Optional[DateTime] = None,
end_datetime: Optional[DateTime] = None,
dropna: bool = True,
) -> Dict[DateTime, Any]:
) -> Dict[str, Any]:
"""Extract a dictionary indexed by the date_time field of the DataRecords.
The dictionary will contain values extracted from the specified key attribute of each DataRecord,
@@ -1130,7 +1146,8 @@ class DataSequence(DataABC, DatabaseRecordProtocolMixin[DataRecord]):
end_timestamp is None or record_date_time_timestamp < end_timestamp
):
filtered_records.append(record)
dates = [record.date_time for record in filtered_records]
# The filter above already excludes records without timestamps.
dates = cast(list[DateTime], [record.date_time for record in filtered_records])
values = [getattr(record, key, None) for record in filtered_records]
return dates, values
@@ -1310,7 +1327,7 @@ class DataSequence(DataABC, DatabaseRecordProtocolMixin[DataRecord]):
query_start = DatabaseTimestamp.to_datetime(query_start_timestamp)
if end_datetime is not None:
# We have a end datetime - look for next entry
end_timestamp = DatabaseTimestamp.from_datetime(query_end)
end_timestamp = DatabaseTimestamp.from_datetime(end_datetime)
query_end_timestamp = await self.db_next_timestamp(end_timestamp)
if query_end_timestamp is None:
# Ensure at least end_datetime is included (excluded by definition)
@@ -1382,13 +1399,12 @@ class DataSequence(DataABC, DatabaseRecordProtocolMixin[DataRecord]):
floored_epoch, unit="s", tz="UTC"
)
else:
resample_origin = resample_start
resample_origin = pd.Timestamp(resample_start)
else:
# Preserve original behaviour: buckets start at the resample start.
resample_origin = resample_start
if resample_origin is None:
# We have no resample origin - take start of day as default
resample_origin = "start_day"
resample_origin = (
pd.Timestamp(resample_start) if resample_start is not None else "start_day"
)
# Check for numeric values
numeric_series = pd.to_numeric(series, errors="coerce") # ensures float64, not object dtype
@@ -1746,7 +1762,7 @@ class DataSequence(DataABC, DatabaseRecordProtocolMixin[DataRecord]):
# ==================== DataProvider ====================
class DataProvider(SingletonMixin, DataSequence):
class DataProvider(SingletonMixin, DataSequence[DataRecordT], Generic[DataRecordT]):
"""Abstract base class for data providers with singleton thread-safety and configurable data parameters.
This class serves as a base for managing generic data, providing an interface for derived
@@ -2013,6 +2029,8 @@ class DataImportMixin(StartMixin):
# Generate value_datetime_mapping once if not using datetime index
if not has_datetime_index:
# Create values datetime list
if start_datetime is None:
raise ValueError("Timezone-aware datetime required")
start_timestamp = DatabaseTimestamp.from_datetime(start_datetime)
value_db_datetimes = list(
self.db_generate_timestamps(start_timestamp, values_count, interval) # type: ignore[attr-defined]
@@ -2094,6 +2112,7 @@ class DataImportMixin(StartMixin):
json_str = json_str.strip() # strip remaining white space at start and end
# Try pandas dataframe with orient="split"
import_data: PydanticDateTimeDataFrame | PydanticDateTimeData | dict[str, Any]
try:
import_data = PydanticDateTimeDataFrame.model_validate_json(json_str)
await self._import_from_dataframe(import_data.to_dataframe())
@@ -2123,7 +2142,7 @@ class DataImportMixin(StartMixin):
# Use simple dict format
try:
import_data = json.loads(json_str)
import_data = cast(dict[str, Any], json.loads(json_str))
await self._import_from_dict(
import_data, key_prefix=key_prefix, start_datetime=start_datetime, interval=interval
)
@@ -2332,7 +2351,7 @@ class DataImportMixin(StartMixin):
# ==================== DataImportProvider ====================
class DataImportProvider(DataImportMixin, DataProvider):
class DataImportProvider(DataImportMixin, DataProvider[DataRecordT], Generic[DataRecordT]):
"""Abstract base class for data providers that import generic data.
This class is designed to handle generic data provided in the form of a key-value dictionary.
@@ -2350,7 +2369,10 @@ class DataImportProvider(DataImportMixin, DataProvider):
# ==================== DataContainer ====================
class DataContainer(SingletonMixin, DataABC):
DataProviderT = TypeVar("DataProviderT", bound=DataProvider)
class DataContainer(SingletonMixin, DataABC, Generic[DataProviderT]):
"""A container for managing multiple DataProvider instances.
This class enables access to data from multiple data providers, supporting retrieval and
@@ -2363,7 +2385,7 @@ class DataContainer(SingletonMixin, DataABC):
"""
# To be overloaded by derived classes.
providers: list[DataProvider] = Field(
providers: list[DataProviderT] = Field(
default_factory=list, json_schema_extra={"description": "List of data providers"}
)
@@ -2381,7 +2403,7 @@ class DataContainer(SingletonMixin, DataABC):
return lock
@field_validator("providers", mode="after")
def check_providers(cls, value: list[DataProvider]) -> list[DataProvider]:
def check_providers(cls, value: list[DataProviderT]) -> list[DataProviderT]:
# Check each item in the list
for item in value:
if not isinstance(item, DataProvider):
@@ -2422,7 +2444,8 @@ class DataContainer(SingletonMixin, DataABC):
return
super().__init__(*args, **kwargs)
def __iter__(self) -> Iterator[str]:
# EOS collections iterate over their elements instead of BaseModel field/value pairs.
def __iter__(self) -> Iterator[str]: # type: ignore[override]
"""Return an iterator over all unique keys available across providers.
Returns:
@@ -2894,7 +2917,7 @@ class DataContainer(SingletonMixin, DataABC):
if key_error:
raise KeyError(f"key `{key}` is not in predictions")
def provider_by_id(self, provider_id: str) -> DataProvider:
def provider_by_id(self, provider_id: str) -> DataProviderT:
"""Retrieves a data provider by its unique identifier.
This method searches through the list of all available providers and
+67 -24
View File
@@ -23,6 +23,7 @@ from typing import (
Type,
TypeVar,
Union,
cast,
)
from loguru import logger
@@ -39,6 +40,7 @@ from akkudoktoreos.core.types import (
ResampleMethod,
)
from akkudoktoreos.utils.datetimeutil import (
UTC,
DateTime,
Duration,
to_datetime,
@@ -278,9 +280,10 @@ class DatabaseBackendABC(ABC, ConfigMixin, SingletonMixin):
class DataRecordProtocol(Protocol):
date_time: DateTime
# Records may be incomplete in memory; database entry points require a timestamp.
date_time: Optional[DateTime]
def __init__(self, date_time: Any) -> None: ...
def __init__(self, date_time: Optional[DateTime]) -> None: ...
def __getitem__(self, key: str) -> Any: ...
@@ -303,15 +306,13 @@ class DatabaseTimestamp(str):
@classmethod
def from_datetime(cls, dt: DateTime) -> "DatabaseTimestamp":
if dt.tz is None:
if dt is None or dt.tz is None:
raise ValueError("Timezone-aware datetime required")
return cls(dt.in_timezone("UTC").format("YYYYMMDDTHHmmss[Z]"))
def to_datetime(self) -> DateTime:
from pendulum import parse
return parse(self)
return to_datetime(self, in_timezone="UTC")
class _DatabaseTimestampUnbound(str):
@@ -801,6 +802,19 @@ class DatabaseRecordProtocolMixin(
return None
def _db_require_date_time(self, record: T_Record) -> DateTime:
"""Validate a record's timestamp before it enters the database index."""
date_time = record.date_time
if date_time is None:
try:
namespace = self.db_namespace()
except NotImplementedError:
namespace = self.__class__.__name__
raise ValueError(
f"Database records require a datetime (namespace='{namespace}', got {record!r})"
)
return date_time
def _db_serialize_record(self, record: T_Record) -> bytes:
"""Serialize a DataRecord to bytes."""
if self.database is None:
@@ -1404,10 +1418,14 @@ class DatabaseRecordProtocolMixin(
if not candidates:
return None
# Indexed records have timestamps, validated when inserted or loaded.
# We validate again to be safe for future refactoring/ changes.
record = min(
candidates,
key=lambda r: abs(
(r.date_time - DatabaseTimestamp.to_datetime(target_timestamp)).total_seconds()
(
self._db_require_date_time(r) - DatabaseTimestamp.to_datetime(target_timestamp)
).total_seconds()
),
)
@@ -1417,7 +1435,8 @@ class DatabaseRecordProtocolMixin(
if (
abs(
(
record.date_time - DatabaseTimestamp.to_datetime(target_timestamp)
self._db_require_date_time(record)
- DatabaseTimestamp.to_datetime(target_timestamp)
).total_seconds()
)
> half_seconds
@@ -1436,7 +1455,7 @@ class DatabaseRecordProtocolMixin(
await self._db_ensure_initialized()
# Ensure normalized to UTC
db_record_date_time = DatabaseTimestamp.from_datetime(record.date_time)
db_record_date_time = DatabaseTimestamp.from_datetime(self._db_require_date_time(record))
await self._db_ensure_loaded(
start_timestamp=db_record_date_time,
@@ -1533,7 +1552,9 @@ class DatabaseRecordProtocolMixin(
continue
record = self._db_deserialize_record(value)
db_record_date_time = DatabaseTimestamp.from_datetime(record.date_time)
db_record_date_time = DatabaseTimestamp.from_datetime(
self._db_require_date_time(record)
)
# Do not resurrect explicitly deleted records
if db_record_date_time in self._db_deleted_timestamps:
@@ -1645,7 +1666,11 @@ class DatabaseRecordProtocolMixin(
start_idx = bisect.bisect_left(self._db_sorted_timestamps, start_timestamp)
for record in self.records[start_idx:]:
record_date_time_timestamp = DatabaseTimestamp.from_datetime(record.date_time)
# Indexed records were validated on insertion or loading.
# We validate againg to be safe for future refactoring/ changes.
record_date_time_timestamp = DatabaseTimestamp.from_datetime(
self._db_require_date_time(record)
)
if start_timestamp and record_date_time_timestamp < start_timestamp:
continue
@@ -1666,7 +1691,9 @@ class DatabaseRecordProtocolMixin(
# Ensure db in memory data and metadata is initialized
await self._db_ensure_initialized()
record_date_time_timestamp = DatabaseTimestamp.from_datetime(record.date_time)
record_date_time_timestamp = DatabaseTimestamp.from_datetime(
self._db_require_date_time(record)
)
self._db_dirty_timestamps.add(record_date_time_timestamp)
# -----------------------------------------------------
@@ -1691,10 +1718,20 @@ class DatabaseRecordProtocolMixin(
save_items = []
for dt in self._db_dirty_timestamps:
record = self._db_record_index.get(dt)
if record:
key = self._db_key_from_timestamp(dt)
value = self._db_serialize_record(record)
save_items.append((key, value))
if record is None:
continue
# Do sanity checks on the record - date_time set and no shift since insertion
current_ts = DatabaseTimestamp.from_datetime(self._db_require_date_time(record))
if current_ts != dt:
raise RuntimeError(
f"Record date_time was mutated after insertion "
f"(index key {dt!r} != current {current_ts!r}); "
"use delete_by_datetime()+insert_by_datetime() to re-time a record."
)
# Add to save
key = self._db_key_from_timestamp(dt)
value = self._db_serialize_record(record)
save_items.append((key, value))
saved_count = len(save_items)
if saved_count:
await self.database.save_records(save_items, namespace=namespace)
@@ -1969,7 +2006,7 @@ class DatabaseRecordProtocolMixin(
# run — they are inside the age window but straddle an incomplete bucket.
raw_cutoff_epoch = int(raw_cutoff_dt.timestamp())
floored_cutoff_epoch = (raw_cutoff_epoch // interval_sec) * interval_sec
new_cutoff_dt = DateTime.fromtimestamp(floored_cutoff_epoch, tz="UTC")
new_cutoff_dt = DateTime.fromtimestamp(floored_cutoff_epoch, tz=UTC)
new_cutoff_ts = DatabaseTimestamp.from_datetime(new_cutoff_dt)
# ---- Determine window start (incremental) ------------------------
@@ -2001,7 +2038,7 @@ class DatabaseRecordProtocolMixin(
# overwritten with the same values).
raw_start_epoch = int(raw_window_start_dt.timestamp())
floored_start_epoch = (raw_start_epoch // interval_sec) * interval_sec
window_start_dt = DateTime.fromtimestamp(floored_start_epoch, tz="UTC")
window_start_dt = DateTime.fromtimestamp(floored_start_epoch, tz=UTC)
window_start_ts = DatabaseTimestamp.from_datetime(window_start_dt)
window_end_dt = new_cutoff_dt # exclusive upper bound, already aligned
@@ -2031,13 +2068,19 @@ class DatabaseRecordProtocolMixin(
# Data is already sparse — check whether timestamps are aligned.
# If every record already sits on an interval boundary, nothing to do.
# If any are misaligned, snap them in place without resampling.
# Indexed records have timestamps, validated when inserted or loaded.
# We validate again to be safe for future refactoring/ changes.
records_in_window = [
r
for r in self.records
if r.date_time is not None and window_start_dt <= r.date_time < window_end_dt
if window_start_dt <= self._db_require_date_time(r) < window_end_dt
]
# The window filter above raises an exception for records without timestamps.
misaligned = [
r for r in records_in_window if int(r.date_time.timestamp()) % interval_sec != 0
r
for r in records_in_window
if int(cast(DateTime, r.date_time).timestamp()) % interval_sec != 0
]
if not misaligned:
logger.debug(
@@ -2065,8 +2108,8 @@ class DatabaseRecordProtocolMixin(
# Process chronologically so the earliest record's values win when
# multiple records floor to the same bucket.
snapped_bucket: dict[int, dict[str, Any]] = {}
for r in sorted(records_in_window, key=lambda x: x.date_time):
ts_epoch = int(r.date_time.timestamp())
for r in sorted(records_in_window, key=lambda r: cast(DateTime, r.date_time)):
ts_epoch = int(cast(DateTime, r.date_time).timestamp())
snapped_epoch = (ts_epoch // interval_sec) * interval_sec
bucket = snapped_bucket.setdefault(snapped_epoch, {})
for key in self.record_keys_writable:
@@ -2089,7 +2132,7 @@ class DatabaseRecordProtocolMixin(
for snapped_epoch, values in snapped_bucket.items():
if not values:
continue
snapped_dt = DateTime.fromtimestamp(snapped_epoch, tz="UTC")
snapped_dt = DateTime.fromtimestamp(snapped_epoch, tz=UTC)
record = self.record_class()(date_time=snapped_dt, **values)
await self.db_insert_record(record, mark_dirty=True)
@@ -2148,7 +2191,7 @@ class DatabaseRecordProtocolMixin(
while first_bucket_epoch < int(window_start_dt.timestamp()):
first_bucket_epoch += interval_sec
compacted_timestamps = [
DateTime.fromtimestamp(first_bucket_epoch + i * interval_sec, tz="UTC")
DateTime.fromtimestamp(first_bucket_epoch + i * interval_sec, tz=UTC)
for i in range(len(array))
]
+1 -1
View File
@@ -92,7 +92,7 @@ class EnergyManagement(
def start_datetime(self) -> DateTime:
"""The starting datetime of the current or latest energy management."""
if EnergyManagement._start_datetime is None:
EnergyManagement.set_start_datetime()
return EnergyManagement.set_start_datetime()
return EnergyManagement._start_datetime
@computed_field # type: ignore[prop-decorator]
+4 -4
View File
@@ -7,7 +7,7 @@ import re
import sys
from pathlib import Path
from types import FrameType
from typing import Any, List, Optional
from typing import Any, List, Optional, cast
import pendulum
from loguru import logger
@@ -176,8 +176,8 @@ def read_file_log(
raise FileNotFoundError("Log file not found")
try:
from_dt = pendulum.parse(from_time) if from_time else None
to_dt = pendulum.parse(to_time) if to_time else None
from_dt = cast(pendulum.DateTime, pendulum.parse(from_time)) if from_time else None
to_dt = cast(pendulum.DateTime, pendulum.parse(to_time)) if to_time else None
except Exception as e:
raise ValueError(f"Invalid date/time format: {e}")
@@ -192,7 +192,7 @@ def read_file_log(
return False
if from_dt or to_dt:
try:
log_time = pendulum.parse(log["time"])
log_time = cast(pendulum.DateTime, pendulum.parse(log["time"]))
except Exception:
return False
if from_dt and log_time < from_dt:
+3 -3
View File
@@ -19,7 +19,7 @@ class LoggingCommonSettings(SettingsBaseModel):
default=None,
json_schema_extra={
"description": "Logging level for API response.",
"examples": LOGGING_LEVELS,
"examples": [*LOGGING_LEVELS],
},
)
@@ -27,7 +27,7 @@ class LoggingCommonSettings(SettingsBaseModel):
default=None,
json_schema_extra={
"description": "Logging level for logging to console.",
"examples": LOGGING_LEVELS,
"examples": [*LOGGING_LEVELS],
},
)
@@ -35,7 +35,7 @@ class LoggingCommonSettings(SettingsBaseModel):
default=None,
json_schema_extra={
"description": "Logging level for logging to file.",
"examples": LOGGING_LEVELS,
"examples": [*LOGGING_LEVELS],
},
)
+30 -15
View File
@@ -20,13 +20,17 @@ import uuid
import weakref
from copy import deepcopy
from typing import (
Annotated,
Any,
Callable,
Dict,
List,
Optional,
Self,
Type,
TypeVar,
Union,
cast,
get_args,
get_origin,
)
@@ -39,6 +43,7 @@ from pydantic import (
BaseModel,
ConfigDict,
Field,
GetPydanticSchema,
PrivateAttr,
RootModel,
ValidationError,
@@ -49,6 +54,7 @@ from pydantic.fields import ComputedFieldInfo, FieldInfo
from akkudoktoreos.utils.datetimeutil import (
DateTime,
Duration,
to_datetime,
to_duration,
to_timezone,
@@ -415,7 +421,7 @@ class PydanticModelNestedValueMixin:
# If this is the final key, set the value
if is_final_key:
try:
model.validate_and_set(key, value)
getattr(model, "validate_and_set")(key, value)
except Exception as e:
raise ValueError(f"Error updating model: {e}") from e
return
@@ -549,10 +555,10 @@ class PydanticModelNestedValueMixin:
if not inspect.isclass(model):
raise TypeError(f"Model '{model}' is not of class type.")
if key not in model.model_fields: # type: ignore[attr-defined]
if key not in model.model_fields:
raise TypeError(f"Field '{key}' does not exist in model '{model.__name__}'.")
field_annotation = model.model_fields[key].annotation # type: ignore[attr-defined]
field_annotation = model.model_fields[key].annotation
if not field_annotation:
raise TypeError(
f"Missing type annotation for field '{key}' in model '{model.__name__}'."
@@ -563,6 +569,8 @@ class PydanticModelNestedValueMixin:
while queue:
annotation = queue.pop(0)
if isinstance(annotation, TypeVar):
annotation = annotation.__bound__ or Any
origin = get_origin(annotation)
args = get_args(annotation)
@@ -679,7 +687,7 @@ class PydanticBaseModel(PydanticModelNestedValueMixin, BaseModel):
"""Resets the fields to their default values."""
for field_name, field_info in self.__class__.model_fields.items():
if field_info.default_factory is not None: # Handle fields with default_factory
default_value = field_info.default_factory()
default_value = field_info.get_default(call_default_factory=True)
else:
default_value = field_info.default
try:
@@ -707,7 +715,7 @@ class PydanticBaseModel(PydanticModelNestedValueMixin, BaseModel):
return self.model_dump()
@classmethod
def from_dict(cls: Type["PydanticBaseModel"], data: dict) -> "PydanticBaseModel":
def from_dict(cls, data: dict) -> Self:
"""Create a PydanticBaseModel instance from a dictionary.
Args:
@@ -735,7 +743,7 @@ class PydanticBaseModel(PydanticModelNestedValueMixin, BaseModel):
return self.model_dump_json()
@classmethod
def from_json(cls: Type["PydanticBaseModel"], json_str: str) -> "PydanticBaseModel":
def from_json(cls, json_str: str) -> Self:
"""Create an instance of the PydanticBaseModel class or its subclass from a JSON string.
Args:
@@ -926,6 +934,10 @@ class PydanticBaseModel(PydanticModelNestedValueMixin, BaseModel):
return None
DateTimeDataInput = dict[str, str | list[float | int | str | None]]
DateTimeDataValues = dict[str, str | DateTime | Duration | list[float | int | str | None]]
class PydanticDateTimeData(RootModel):
"""Pydantic model for time series data with consistent value lengths.
@@ -948,13 +960,16 @@ class PydanticDateTimeData(RootModel):
"""
root: Dict[str, Union[str, List[Union[float, int, str, None]]]]
# The wire format contains strings; validate_root normalizes the two
# indexing values to Pendulum objects. Keep the existing input schema.
root: Annotated[
DateTimeDataValues,
GetPydanticSchema(lambda source_type, handler: handler(DateTimeDataInput)),
]
@field_validator("root", mode="after")
@classmethod
def validate_root(
cls, value: Dict[str, Union[str, List[Union[float, int, str, None]]]]
) -> Dict[str, Union[str, List[Union[float, int, str, None]]]]:
def validate_root(cls, value: dict[str, Any]) -> DateTimeDataValues:
# Validate that all keys are strings
if not all(isinstance(k, str) for k in value.keys()):
raise ValueError("All keys in the dictionary must be strings.")
@@ -977,7 +992,7 @@ class PydanticDateTimeData(RootModel):
return value
def to_dict(self) -> Dict[str, Union[str, List[Union[float, int, str, None]]]]:
def to_dict(self) -> DateTimeDataValues:
"""Convert the model to a plain dictionary.
Returns:
@@ -1176,8 +1191,8 @@ class PydanticDateTimeDataFrame(PydanticBaseModel):
df[col] = df[col].dt.tz_convert(resolved_tz)
return cls(
data=df.to_dict(orient="index"),
dtypes={col: str(dtype) for col, dtype in df.dtypes.items()},
data=cast(dict[str, dict[str, Any]], df.to_dict(orient="index")),
dtypes=cast(dict[str, str], {col: str(dtype) for col, dtype in df.dtypes.items()}),
tz=resolved_tz,
datetime_columns=datetime_columns,
)
@@ -1401,10 +1416,10 @@ class PydanticDateTimeSeries(PydanticBaseModel):
series.index = index
if len(index) > 0:
tz = to_datetime(series.index[0]).timezone.name
tz = to_datetime(series.index[0]).timezone_name
return cls(
data=series.to_dict(),
data=cast(dict[str, Any], series.to_dict()),
dtype=str(series.dtype),
tz=tz,
)
+3 -3
View File
@@ -106,7 +106,7 @@ class BatteriesCommonSettings(DevicesBaseSettings):
def validate_and_sort_charge_rates(cls, v: Any) -> NDArray[Shape["*"], float]:
# None means fallback to default values
if v is None:
return BATTERY_DEFAULT_CHARGE_RATES.copy()
return np.asarray(BATTERY_DEFAULT_CHARGE_RATES, dtype=float)
# Convert to numpy array
if isinstance(v, str):
@@ -345,10 +345,10 @@ class DevicesCommonSettings(SettingsBaseModel):
if self.max_batteries and self.batteries:
for battery in self.batteries:
keys.extend(battery.measurement_keys)
keys.extend(battery.measurement_keys or [])
if self.max_electric_vehicles and self.electric_vehicles:
for electric_vehicle in self.electric_vehicles:
keys.extend(electric_vehicle.measurement_keys)
keys.extend(electric_vehicle.measurement_keys or [])
return keys
+5 -1
View File
@@ -94,7 +94,7 @@ class MeasurementDataRecord(DataRecord):
return keys
class Measurement(SingletonMixin, DataImportMixin, DataSequence):
class Measurement(SingletonMixin, DataImportMixin, DataSequence[MeasurementDataRecord]):
"""Singleton class that holds measurement data records.
Measurements can be provided programmatically or read from JSON string or file.
@@ -168,6 +168,8 @@ class Measurement(SingletonMixin, DataImportMixin, DataSequence):
np.ndarray: A NumPy Array of the energy [kWh] per interval values calculated from
the meter readings.
"""
if start_datetime is None or end_datetime is None:
raise ValueError("Start and end datetimes are required for energy calculation")
size = self._interval_count(start_datetime, end_datetime, interval)
energy_mr_array = await self.key_to_array(
@@ -237,6 +239,8 @@ class Measurement(SingletonMixin, DataImportMixin, DataSequence):
end_datetime = await self.max_datetime()
if end_datetime:
end_datetime = end_datetime.add(seconds=1)
if start_datetime is None or end_datetime is None:
raise ValueError("Start and end datetimes are required for energy calculation")
size = self._interval_count(start_datetime, end_datetime, interval)
load_total_kwh_array = np.zeros(size)
@@ -125,7 +125,7 @@ class GeneticSimulation(PydanticBaseModel):
self.pv_prediction_wh = np.array(parameters.pv_forecast_wh, float)
self.elect_price_hourly = np.array(parameters.electricity_price_per_wh, float)
self.elect_revenue_per_hour_arr = (
parameters.feed_in_tariff_per_wh
np.asarray(parameters.feed_in_tariff_per_wh, dtype=float)
if isinstance(parameters.feed_in_tariff_per_wh, list)
else np.full(len(self.load_energy_array), parameters.feed_in_tariff_per_wh, float)
)
@@ -1205,21 +1205,21 @@ class GeneticOptimization(OptimizationBase):
)
# Simulation may have changed something, use simulation values
ac_charge_hours = self.simulation.ac_charge_hours
if ac_charge_hours is None:
ac_charge_hours = []
else:
ac_charge_hours = ac_charge_hours.tolist()
dc_charge_hours = self.simulation.dc_charge_hours
if dc_charge_hours is None:
dc_charge_hours = []
else:
dc_charge_hours = dc_charge_hours.tolist()
discharge = self.simulation.bat_discharge_hours
if discharge is None:
discharge = []
else:
discharge = discharge.tolist()
ac_charge_hours = (
self.simulation.ac_charge_hours.tolist()
if self.simulation.ac_charge_hours is not None
else []
)
dc_charge_hours = (
self.simulation.dc_charge_hours.tolist()
if self.simulation.dc_charge_hours is not None
else []
)
discharge = (
self.simulation.bat_discharge_hours.tolist()
if self.simulation.bat_discharge_hours is not None
else []
)
return GeneticSolution(
**{
@@ -62,10 +62,10 @@ class GeneticCommonSettings(SettingsBaseModel):
# --- Penalties (existing) -------------------------------------------------
penalties: dict[str, Union[float, int, str]] = Field(
default_factory=lambda: {
"ev_soc_miss": 10,
"ac_charge_break_even": 1.0,
},
default_factory=lambda: dict[str, float | int | str](
ev_soc_miss=10,
ac_charge_break_even=1.0,
),
json_schema_extra={
"description": "Penalty parameters used in fitness evaluation.",
"examples": [{"ev_soc_miss": 10}],
@@ -141,7 +141,11 @@ class GeneticVisualizationReport(ConfigMixin):
marker = markers[idx] if markers and idx < len(markers) else "o" # Marker style
line_style = line_styles[idx] if line_styles and idx < len(line_styles) else "-"
plt.plot(
timestamps, y_data, label=label, marker=marker, linestyle=line_style
mdates.date2num(timestamps),
np.asarray(y_data, dtype=float),
label=label,
marker=marker,
linestyle=line_style,
) # Plot line
# Format the time axis
@@ -178,8 +182,15 @@ class GeneticVisualizationReport(ConfigMixin):
# Add vertical line for the current date if within the axis range
current_time = pendulum.now(self.config.general.timezone)
if timestamps[0].subtract(hours=2) <= current_time <= timestamps[-1]:
plt.axvline(current_time, color="r", linestyle="--", label="Now")
plt.text(current_time, plt.ylim()[1], "Now", color="r", ha="center", va="bottom")
plt.axvline(mdates.date2num(current_time), color="r", linestyle="--", label="Now")
plt.text(
mdates.date2num(current_time),
plt.ylim()[1],
"Now",
color="r",
ha="center",
va="bottom",
)
# Add a second x-axis on top
ax1 = plt.gca()
@@ -191,7 +202,9 @@ class GeneticVisualizationReport(ConfigMixin):
# ax2.set_xticks(timestamps[::48]) # Set ticks every 12 hours
# ax2.set_xticklabels([f"{int(h)}" for h in hours_since_start[::48]])
# ax2.set_xticks(timestamps[:: len(timestamps) // 24]) # Select 10 evenly spaced ticks
ax2.set_xticks(timestamps[:: len(timestamps) // 12]) # Select 10 evenly spaced ticks
ax2.set_xticks(
mdates.date2num(timestamps[:: len(timestamps) // 12])
) # Select 10 evenly spaced ticks
# ax2.set_xticklabels([f"{int(h)}" for h in hours_since_start[:: len(timestamps) // 24]])
ax2.set_xticklabels([f"{int(h)}" for h in hours_since_start[:: len(timestamps) // 12]])
if x2label:
@@ -251,7 +264,13 @@ class GeneticVisualizationReport(ConfigMixin):
line_style = (
line_styles[idx] if line_styles and idx < len(line_styles) else "-"
) # Line style
plt.plot(x, y_data, label=label, marker=marker, linestyle=line_style) # Plot line
plt.plot(
x,
np.asarray(y_data, dtype=float),
label=label,
marker=marker,
linestyle=line_style,
) # Plot line
plt.title(title) # Set title
plt.xlabel(xlabel) # Set x-axis label
@@ -16,6 +16,7 @@ from akkudoktoreos.devices.genetic0.genetic0homeappliance import Genetic0HomeApp
from akkudoktoreos.devices.genetic0.genetic0inverter import Genetic0Inverter
from akkudoktoreos.optimization.genetic0.genetic0params import (
Genetic0EnergyManagementParameters,
Genetic0OptimizationParameters,
)
from akkudoktoreos.optimization.genetic0.genetic0solution import (
Genetic0SimulationResult,
@@ -128,7 +129,7 @@ class Genetic0Simulation(PydanticBaseModel):
self.pv_prediction_wh = np.array(parameters.pv_forecast_wh, float)
self.elect_price_hourly = np.array(parameters.electricity_price_per_wh, float)
self.elect_revenue_per_hour_arr = (
parameters.feed_in_tariff_per_wh
np.asarray(parameters.feed_in_tariff_per_wh, dtype=float)
if isinstance(parameters.feed_in_tariff_per_wh, list)
else np.full(len(self.load_energy_array), parameters.feed_in_tariff_per_wh, float)
)
@@ -741,7 +742,7 @@ class Genetic0Optimization(OptimizationBase):
def evaluate(
self,
individual: list[int],
parameters: Genetic0EnergyManagementParameters,
parameters: Genetic0OptimizationParameters,
start_hour: int,
worst_case: bool,
) -> tuple[float]:
@@ -1056,7 +1057,7 @@ class Genetic0Optimization(OptimizationBase):
def optimize_ems(
self,
parameters: Genetic0EnergyManagementParameters,
parameters: Genetic0OptimizationParameters,
start_hour: Optional[int] = None,
worst_case: bool = False,
ngen: Optional[int] = None,
@@ -1208,21 +1209,21 @@ class Genetic0Optimization(OptimizationBase):
)
# Simulation may have changed something, use simulation values
ac_charge_hours = self.simulation.ac_charge_hours
if ac_charge_hours is None:
ac_charge_hours = []
else:
ac_charge_hours = ac_charge_hours.tolist()
dc_charge_hours = self.simulation.dc_charge_hours
if dc_charge_hours is None:
dc_charge_hours = []
else:
dc_charge_hours = dc_charge_hours.tolist()
discharge = self.simulation.bat_discharge_hours
if discharge is None:
discharge = []
else:
discharge = discharge.tolist()
ac_charge_hours = (
self.simulation.ac_charge_hours.tolist()
if self.simulation.ac_charge_hours is not None
else []
)
dc_charge_hours = (
self.simulation.dc_charge_hours.tolist()
if self.simulation.dc_charge_hours is not None
else []
)
discharge = (
self.simulation.bat_discharge_hours.tolist()
if self.simulation.bat_discharge_hours is not None
else []
)
return Genetic0Solution(
**{
@@ -52,10 +52,10 @@ class Genetic0CommonSettings(SettingsBaseModel):
# --- Penalties (existing) -------------------------------------------------
penalties: dict[str, Union[float, int, str]] = Field(
default_factory=lambda: {
"ev_soc_miss": 10,
"ac_charge_break_even": 1.0,
},
default_factory=lambda: dict[str, float | int | str](
ev_soc_miss=10,
ac_charge_break_even=1.0,
),
json_schema_extra={
"description": "Penalty parameters used in fitness evaluation.",
"examples": [{"ev_soc_miss": 10}],
@@ -141,7 +141,11 @@ class Genetic0VisualizationReport(ConfigMixin):
marker = markers[idx] if markers and idx < len(markers) else "o" # Marker style
line_style = line_styles[idx] if line_styles and idx < len(line_styles) else "-"
plt.plot(
timestamps, y_data, label=label, marker=marker, linestyle=line_style
mdates.date2num(timestamps),
np.asarray(y_data, dtype=float),
label=label,
marker=marker,
linestyle=line_style,
) # Plot line
# Format the time axis
@@ -178,8 +182,15 @@ class Genetic0VisualizationReport(ConfigMixin):
# Add vertical line for the current date if within the axis range
current_time = pendulum.now(self.config.general.timezone)
if timestamps[0].subtract(hours=2) <= current_time <= timestamps[-1]:
plt.axvline(current_time, color="r", linestyle="--", label="Now")
plt.text(current_time, plt.ylim()[1], "Now", color="r", ha="center", va="bottom")
plt.axvline(mdates.date2num(current_time), color="r", linestyle="--", label="Now")
plt.text(
mdates.date2num(current_time),
plt.ylim()[1],
"Now",
color="r",
ha="center",
va="bottom",
)
# Add a second x-axis on top
ax1 = plt.gca()
@@ -191,7 +202,9 @@ class Genetic0VisualizationReport(ConfigMixin):
# ax2.set_xticks(timestamps[::48]) # Set ticks every 12 hours
# ax2.set_xticklabels([f"{int(h)}" for h in hours_since_start[::48]])
# ax2.set_xticks(timestamps[:: len(timestamps) // 24]) # Select 10 evenly spaced ticks
ax2.set_xticks(timestamps[:: len(timestamps) // 12]) # Select 10 evenly spaced ticks
ax2.set_xticks(
mdates.date2num(timestamps[:: len(timestamps) // 12])
) # Select 10 evenly spaced ticks
# ax2.set_xticklabels([f"{int(h)}" for h in hours_since_start[:: len(timestamps) // 24]])
ax2.set_xticklabels([f"{int(h)}" for h in hours_since_start[:: len(timestamps) // 12]])
if x2label:
@@ -251,7 +264,13 @@ class Genetic0VisualizationReport(ConfigMixin):
line_style = (
line_styles[idx] if line_styles and idx < len(line_styles) else "-"
) # Line style
plt.plot(x, y_data, label=label, marker=marker, linestyle=line_style) # Plot line
plt.plot(
x,
np.asarray(y_data, dtype=float),
label=label,
marker=marker,
linestyle=line_style,
) # Plot line
plt.title(title) # Set title
plt.xlabel(xlabel) # Set x-axis label
+1 -1
View File
@@ -104,7 +104,7 @@ class ElecFeeDataRecord(PredictionRecord):
return self.elecfee_feedin_amt_wh * 1000.0
class ElecFeeProvider(PredictionProvider):
class ElecFeeProvider(PredictionProvider[ElecFeeDataRecord]):
"""Abstract base class for electricity fee providers.
Electricity fee providers predict fees on consumed and feed-in electricity to be used by
+1 -1
View File
@@ -49,7 +49,7 @@ class ElecPriceDataRecord(PredictionRecord):
return self.elecprice_marketprice_wh * 1000.0
class ElecPriceProvider(PricePredictionProviderBase):
class ElecPriceProvider(PricePredictionProviderBase[ElecPriceDataRecord]):
"""Abstract base class for electricity price providers.
ElecPriceProvider is a thread-safe singleton, ensuring only one instance of this class is created.
@@ -49,7 +49,7 @@ class FeedInTariffDataRecord(PredictionRecord):
return self.feed_in_tariff_wh * 1000.0
class FeedInTariffProvider(PricePredictionProviderBase):
class FeedInTariffProvider(PricePredictionProviderBase[FeedInTariffDataRecord]):
"""Abstract base class for feed in tariff providers.
FeedInTariffProvider is a thread-safe singleton, ensuring only one instance of this class is created.
@@ -123,10 +123,10 @@ class FeedInTariffAkkudoktor(FeedInTariffProvider):
history = np.asarray(
await self.key_to_array(
key="feed_in_tariff_wh",
end_datetime=self.highest_orig_datetime,
end_datetime=to_datetime(self.highest_orig_datetime),
fill_method="linear",
),
dtype=float,
dtype=np.float64,
)
covered_hours = (
int((self.highest_orig_datetime - self.ems_start_datetime).total_seconds() // 3600) + 1
@@ -279,7 +279,7 @@ class FeedInTariffEnergyCharts(FeedInTariffProvider):
# above, so ETS/median always trains on the true wholesale-price signal.
history = await self.key_to_array(
key="feed_in_tariff_raw_wh",
end_datetime=self.highest_orig_datetime,
end_datetime=to_datetime(self.highest_orig_datetime),
interval=to_duration(f"{resolution_seconds} seconds"),
fill_method="linear",
)
@@ -137,11 +137,11 @@ class FeedInTariffTibber(FeedInTariffProvider):
history = np.asarray(
await self.key_to_array(
key="feed_in_tariff_wh",
end_datetime=self.highest_orig_datetime,
end_datetime=to_datetime(self.highest_orig_datetime),
interval=to_duration(f"{interval_seconds} seconds"),
fill_method="linear",
),
dtype=float,
dtype=np.float64,
)
covered_slots = 0
if self.highest_orig_datetime >= self.ems_start_datetime:
+6 -3
View File
@@ -5,7 +5,7 @@ Notes:
"""
from abc import abstractmethod
from typing import List, Optional
from typing import Generic, List, Optional, TypeVar
from pydantic import Field
@@ -20,7 +20,10 @@ class LoadDataRecord(PredictionRecord):
)
class LoadProvider(PredictionProvider):
LoadDataRecordT = TypeVar("LoadDataRecordT", bound=LoadDataRecord)
class LoadProvider(PredictionProvider[LoadDataRecordT], Generic[LoadDataRecordT]):
"""Abstract base class for load providers.
LoadProvider is a thread-safe singleton, ensuring only one instance of this class is created.
@@ -41,7 +44,7 @@ class LoadProvider(PredictionProvider):
"""
# overload
records: List[LoadDataRecord] = Field(
records: List[LoadDataRecordT] = Field(
default_factory=list, json_schema_extra={"description": "List of LoadDataRecord records"}
)
@@ -32,7 +32,7 @@ class LoadAkkudoktorDataRecord(LoadDataRecord):
)
class LoadAkkudoktor(LoadProvider):
class LoadAkkudoktor(LoadProvider[LoadAkkudoktorDataRecord]):
"""Fetch Load forecast data from Akkudoktor load profiles."""
records: list[LoadAkkudoktorDataRecord] = Field(
+38 -70
View File
@@ -119,40 +119,42 @@ weather_openmeteo = WeatherOpenMeteo()
weather_import = WeatherImport()
def prediction_providers() -> list[
Union[
ElecFeeFixed,
ElecFeeImport,
ElecPriceAkkudoktor,
ElecPriceEnergyCharts,
ElecPriceFixed,
ElecPriceImport,
ElecPriceSMARD,
ElecPriceTibber,
FeedInTariffAkkudoktor,
FeedInTariffDvhubOnline,
FeedInTariffEnergyCharts,
FeedInTariffFixed,
FeedInTariffImport,
FeedInTariffSMARD,
FeedInTariffTibber,
LoadAkkudoktor,
LoadAkkudoktorAdjusted,
LoadImport,
LoadVrm,
PVForecastAkkudoktor,
PVForecastForecastSolar,
PVForecastImport,
PVForecastPVLib,
PVForecastPVNode,
PVForecastSolcast,
PVForecastVrm,
WeatherBrightSky,
WeatherClearOutside,
WeatherImport,
WeatherOpenMeteo,
]
]:
PredictionProviderType = Union[
ElecFeeFixed,
ElecFeeImport,
ElecPriceAkkudoktor,
ElecPriceEnergyCharts,
ElecPriceFixed,
ElecPriceImport,
ElecPriceSMARD,
ElecPriceTibber,
FeedInTariffAkkudoktor,
FeedInTariffDvhubOnline,
FeedInTariffEnergyCharts,
FeedInTariffFixed,
FeedInTariffImport,
FeedInTariffSMARD,
FeedInTariffTibber,
LoadAkkudoktor,
LoadAkkudoktorAdjusted,
LoadImport,
LoadVrm,
PVForecastAkkudoktor,
PVForecastForecastSolar,
PVForecastHomeAssistant,
PVForecastImport,
PVForecastPVLib,
PVForecastPVNode,
PVForecastSolcast,
PVForecastVrm,
WeatherBrightSky,
WeatherClearOutside,
WeatherImport,
WeatherOpenMeteo,
]
def prediction_providers() -> list[PredictionProviderType]:
"""Return list of prediction providers.
Factory for prediction container.
@@ -229,44 +231,10 @@ def prediction_providers() -> list[
]
class Prediction(PredictionContainer):
class Prediction(PredictionContainer[PredictionProviderType]):
"""Prediction container to manage multiple prediction providers."""
providers: list[
Union[
ElecFeeFixed,
ElecFeeImport,
ElecPriceAkkudoktor,
ElecPriceEnergyCharts,
ElecPriceFixed,
ElecPriceImport,
ElecPriceSMARD,
ElecPriceTibber,
FeedInTariffAkkudoktor,
FeedInTariffDvhubOnline,
FeedInTariffEnergyCharts,
FeedInTariffFixed,
FeedInTariffImport,
FeedInTariffSMARD,
FeedInTariffTibber,
LoadAkkudoktor,
LoadAkkudoktorAdjusted,
LoadImport,
LoadVrm,
PVForecastAkkudoktor,
PVForecastForecastSolar,
PVForecastHomeAssistant,
PVForecastImport,
PVForecastPVLib,
PVForecastPVNode,
PVForecastSolcast,
PVForecastVrm,
WeatherBrightSky,
WeatherClearOutside,
WeatherImport,
WeatherOpenMeteo,
]
] = Field(
providers: list[PredictionProviderType] = Field(
default_factory=prediction_providers,
json_schema_extra={"description": "List of prediction providers"},
)
+20 -7
View File
@@ -8,7 +8,7 @@ This module is designed for use in predictive modeling workflows, facilitating t
and manipulation of configuration and prediction data in a clear, scalable, and structured manner.
"""
from typing import List, Optional
from typing import Generic, List, Optional, TypeVar
from loguru import logger
from pydantic import Field, computed_field
@@ -20,6 +20,7 @@ from akkudoktoreos.core.dataabc import (
DataImportProvider,
DataProvider,
DataRecord,
DataRecordT,
DataSequence,
)
from akkudoktoreos.utils.datetimeutil import DateTime, Duration, to_duration
@@ -52,7 +53,10 @@ class PredictionRecord(DataRecord):
pass
class PredictionSequence(DataSequence):
PredictionRecordT = TypeVar("PredictionRecordT", bound=PredictionRecord)
class PredictionSequence(DataSequence[PredictionRecordT], Generic[PredictionRecordT]):
"""A managed sequence of PredictionRecord instances with list-like behavior.
The PredictionSequence class provides an ordered, mutable collection of PredictionRecord
@@ -90,7 +94,7 @@ class PredictionSequence(DataSequence):
"""
# To be overloaded by derived classes.
records: List[PredictionRecord] = Field(
records: List[PredictionRecordT] = Field(
default_factory=list, json_schema_extra={"description": "List of prediction records"}
)
@@ -185,7 +189,9 @@ class PredictionStartEndKeepMixin(PredictionABC):
return int(duration.total_hours())
class PredictionProvider(PredictionStartEndKeepMixin, DataProvider):
class PredictionProvider(
PredictionStartEndKeepMixin, DataProvider[DataRecordT], Generic[DataRecordT]
):
"""Abstract base class for prediction providers with singleton thread-safety and configurable prediction parameters.
This class serves as a base for managing prediction data, providing an interface for derived
@@ -249,7 +255,9 @@ class PredictionProvider(PredictionStartEndKeepMixin, DataProvider):
await self._update_data(force_update=force_update)
class PredictionImportProvider(PredictionProvider, DataImportProvider):
class PredictionImportProvider(
PredictionProvider[DataRecordT], DataImportProvider[DataRecordT], Generic[DataRecordT]
):
"""Abstract base class for prediction providers that import prediction data.
This class is designed to handle prediction data provided in the form of a key-value dictionary.
@@ -264,7 +272,12 @@ class PredictionImportProvider(PredictionProvider, DataImportProvider):
pass
class PredictionContainer(PredictionStartEndKeepMixin, DataContainer):
PredictionProviderT = TypeVar("PredictionProviderT", bound=PredictionProvider)
class PredictionContainer(
PredictionStartEndKeepMixin, DataContainer[PredictionProviderT], Generic[PredictionProviderT]
):
"""A container for managing multiple PredictionProvider instances.
This class enables access to data from multiple prediction providers, supporting retrieval and
@@ -277,6 +290,6 @@ class PredictionContainer(PredictionStartEndKeepMixin, DataContainer):
"""
# To be overloaded by derived classes.
providers: List[PredictionProvider] = Field(
providers: List[PredictionProviderT] = Field(
default_factory=list, json_schema_extra={"description": "List of prediction providers"}
)
+5 -2
View File
@@ -1,7 +1,7 @@
"""Shared base for price-like predictions (electricity price, feed-in tariff)."""
from abc import abstractmethod
from typing import cast
from typing import Generic, cast
import numpy as np
import pandas as pd
@@ -9,11 +9,14 @@ from loguru import logger
from statsmodels.tsa.holtwinters import ExponentialSmoothing
from akkudoktoreos.core.coreabc import PredictionMixin
from akkudoktoreos.core.dataabc import DataRecordT
from akkudoktoreos.prediction.predictionabc import PredictionProvider
from akkudoktoreos.utils.datetimeutil import DateTime, to_datetime, to_duration
class PricePredictionProviderBase(PredictionMixin, PredictionProvider):
class PricePredictionProviderBase(
PredictionMixin, PredictionProvider[DataRecordT], Generic[DataRecordT]
):
"""Common forecasting + fee-application logic shared by price-like providers.
Subclasses must supply the raw/gross record keys, the fee keys to pull from
@@ -5,7 +5,7 @@ Notes:
"""
from abc import abstractmethod
from typing import List, Optional
from typing import Generic, List, Optional, TypeVar
from loguru import logger
from pydantic import Field
@@ -24,7 +24,10 @@ class PVForecastDataRecord(PredictionRecord):
)
class PVForecastProvider(PredictionProvider):
PVForecastDataRecordT = TypeVar("PVForecastDataRecordT", bound=PVForecastDataRecord)
class PVForecastProvider(PredictionProvider[PVForecastDataRecordT], Generic[PVForecastDataRecordT]):
"""Abstract base class for pvforecast providers.
PVForecastProvider is a thread-safe singleton, ensuring only one instance of this class is created.
@@ -45,7 +48,7 @@ class PVForecastProvider(PredictionProvider):
"""
# overload
records: List[PVForecastDataRecord] = Field(
records: List[PVForecastDataRecordT] = Field(
default_factory=list,
json_schema_extra={"description": "List of PVForecastDataRecord records"},
)
@@ -188,7 +188,7 @@ class PVForecastAkkudoktorDataRecord(PVForecastDataRecord):
return self.pvforecast_ac_power
class PVForecastAkkudoktor(PVForecastProvider):
class PVForecastAkkudoktor(PVForecastProvider[PVForecastAkkudoktorDataRecord]):
"""Fetch and process PV forecast data from akkudoktor.net.
PVForecastAkkudoktor is a singleton-based class that retrieves weather forecast data
@@ -72,6 +72,8 @@ class PVForecastForecastSolar(PVForecastProvider):
return to_datetime(s)
tz = iana_tz or str(self.config.general.timezone)
dt = pendulum.parse(s, tz=tz)
if not isinstance(dt, pendulum.DateTime):
raise ValueError(f"Expected a datetime, got {local_ts!r}")
return to_datetime(dt.isoformat())
def _plane_url(self, plane: Any) -> str:
@@ -100,6 +100,8 @@ class PVForecastPVNode(PVForecastProvider):
tz = iana_tz or str(self.config.general.timezone)
# Interpret the naive wall-clock string AS local time in tz, then resolve.
dt = pendulum.parse(s, tz=tz)
if not isinstance(dt, pendulum.DateTime):
raise ValueError(f"Expected a datetime, got {local_ts!r}")
return to_datetime(dt.isoformat())
def _extract_values(self, body: Any) -> list[tuple[Any, float]]:
+1 -1
View File
@@ -112,7 +112,7 @@ class WeatherDataRecord(PredictionRecord):
)
class WeatherProvider(PredictionProvider):
class WeatherProvider(PredictionProvider[WeatherDataRecord]):
"""Abstract base class for weather providers.
WeatherProvider is a thread-safe singleton, ensuring only one instance of this class is created.
@@ -246,12 +246,15 @@ class WeatherBrightSky(WeatherProvider):
logger.debug(debug_msg)
return
data = pvlib.atmosphere.gueymard94_pw(temperature, humidity)
end_datetime = self.end_datetime
if end_datetime is None:
raise ValueError("Prediction end datetime is not available")
pwat = pd.Series(
data=data,
index=pd.DatetimeIndex(
pd.date_range(
start=self.ems_start_datetime,
end=self.end_datetime,
end=end_datetime,
freq="1h",
inclusive="left",
)
@@ -13,7 +13,7 @@ Notes:
"""
import re
from typing import Dict, List, Optional, Tuple
from typing import Any, Dict, List, Optional, Tuple
import pandas as pd
import requests
@@ -232,7 +232,7 @@ class WeatherClearOutside(WeatherProvider):
p_detail_tables.pop(0)
# Create clearout data
clearout_data = {}
clearout_data: dict[str, Any] = {}
# Number of detail values. On last day may be less than 24.
detail_values_count = None
# Add data values
@@ -258,7 +258,7 @@ class WeatherClearOutside(WeatherProvider):
raise ValueError(error_msg)
# Scrape the detail values
detail_data = []
detail_data: list[float | str] = []
extra_detail_name = None
extra_detail_data = []
for p_detail_value in p_detail_values:
@@ -281,9 +281,10 @@ class WeatherClearOutside(WeatherProvider):
and hasattr(p_detail_value, "title")
and p_detail_value.title
):
value_str = p_detail_value.title.string
value_str = p_detail_value.title.get_text()
else:
value_str = p_detail_value.get_text()
value: float | str
try:
value = float(value_str)
except ValueError:
@@ -336,9 +337,9 @@ class WeatherClearOutside(WeatherProvider):
if key is None:
continue
if detail_name in clearout_data:
value = clearout_data[detail_name][row_index]
record_value = clearout_data[detail_name][row_index]
corr_factor = clearoutside_key_mapping[detail_name][1]
if corr_factor:
value = value * corr_factor
setattr(weather_record, key, value)
record_value = record_value * corr_factor
setattr(weather_record, key, record_value)
await self.insert_by_datetime(weather_record)
@@ -342,12 +342,15 @@ class WeatherOpenMeteo(WeatherProvider):
return
data = pvlib.atmosphere.gueymard94_pw(temperature, humidity)
end_datetime = self.end_datetime
if end_datetime is None:
raise ValueError("Prediction end datetime is not available")
pwat = pd.Series(
data=data,
index=pd.DatetimeIndex(
pd.date_range(
start=self.ems_start_datetime,
end=self.end_datetime,
end=end_datetime,
freq="1h",
inclusive="left",
)
+9 -3
View File
@@ -1,12 +1,18 @@
# Module taken from https://github.com/koaning/fh-altair
# MIT license
from typing import Optional
from typing import Callable, Optional, cast
from bokeh.embed import components
from bokeh.models import Plot
from bokeh.models.annotations import Title
from bokeh.plotting import figure
from bokeh.resources import INLINE
from monsterui.franken import H4, Card, NotStr
# Bokeh accepts FigureOptions as constructor keywords, but its generated
# constructor signature only lists model properties. Preserve the typed result.
create_figure = cast(Callable[..., figure], figure)
# Javascript for bokeh - to be included by the page
BokehJS = [NotStr(INLINE.render_css()), NotStr(INLINE.render_js())]
@@ -29,7 +35,7 @@ def bokey_apply_theme_to_plot(plot: Plot, dark: bool) -> None:
if dark:
plot.background_fill_color = "#1e1e1e"
plot.border_fill_color = "#1e1e1e"
plot.title.text_color = "white"
cast(Title, plot.title).text_color = "white"
for ax in plot.xaxis + plot.yaxis:
ax.axis_line_color = "white"
ax.major_tick_line_color = "white"
@@ -44,7 +50,7 @@ def bokey_apply_theme_to_plot(plot: Plot, dark: bool) -> None:
else:
plot.background_fill_color = "white"
plot.border_fill_color = "white"
plot.title.text_color = "black"
cast(Title, plot.title).text_color = "black"
for ax in plot.xaxis + plot.yaxis:
ax.axis_line_color = "black"
ax.major_tick_line_color = "black"
+1 -1
View File
@@ -59,7 +59,7 @@ def item_model_defaults(item_model: Any) -> tuple[dict, list[str]]:
if field_info.default is not PydanticUndefined:
kwargs[field_name] = field_info.default
elif field_info.default_factory is not None:
kwargs[field_name] = field_info.default_factory()
kwargs[field_name] = field_info.get_default(call_default_factory=True)
else:
required_missing.append(field_name)
@@ -192,7 +192,7 @@ def get_default_value(field_info: Union[FieldInfo, ComputedFieldInfo], regular_f
"""
import pathlib
if not regular_field:
if not regular_field or not isinstance(field_info, FieldInfo):
return "N/A"
# Resolve the raw default — prefer plain default, fall back to factory
@@ -200,7 +200,7 @@ def get_default_value(field_info: Union[FieldInfo, ComputedFieldInfo], regular_f
val = field_info.default
elif field_info.default_factory is not None:
try:
val = field_info.default_factory()
val = field_info.get_default(call_default_factory=True)
except Exception:
return ""
else:
@@ -250,12 +250,12 @@ def resolve_nested_types(field_type: Any, parent_types: list[str]) -> list[tuple
def create_config_details(
model: type[PydanticBaseModel], values: dict, values_prefix: list[str] = []
model: type[PydanticBaseModel] | type[ConfigEOS], values: dict, values_prefix: list[str] = []
) -> dict[str, dict]:
"""Generate configuration details based on provided values and model metadata.
Args:
model (type[PydanticBaseModel]): The Pydantic model to extract configuration from.
model: An EOS model or the top-level settings class to extract configuration from.
values (dict): A dictionary containing the current configuration values.
values_prefix (list[str]): A list of parent type names that prefixes the model values in the values.
@@ -271,7 +271,11 @@ def create_config_details(
) -> None:
nonlocal values, values_prefix
regular_field = isinstance(subfield_info, FieldInfo)
subtype = subfield_info.annotation if regular_field else subfield_info.return_type
subtype = (
subfield_info.annotation
if isinstance(subfield_info, FieldInfo)
else subfield_info.return_type
)
nested_types = resolve_nested_types(subtype, [])
found_basic = False
+13 -8
View File
@@ -9,6 +9,7 @@ from fasthtml.common import FT, Div, NotStr
from markdown_it import MarkdownIt
from markdown_it.renderer import RendererHTML
from markdown_it.token import Token
from markdown_it.utils import OptionsDict
from monsterui.foundations import stringify
# Where to find the static data assets
@@ -42,7 +43,7 @@ def file_to_data_uri(file_path: Path) -> str:
def render_heading(
self: RendererHTML, tokens: List[Token], idx: int, options: dict, env: dict
self: RendererHTML, tokens: List[Token], idx: int, options: OptionsDict, env: dict
) -> str:
"""Custom renderer for Markdown headings with MonsterUI styling."""
if tokens[idx].markup == "#":
@@ -63,7 +64,7 @@ def render_heading(
def render_paragraph(
self: RendererHTML, tokens: List[Token], idx: int, options: dict, env: dict
self: RendererHTML, tokens: List[Token], idx: int, options: OptionsDict, env: dict
) -> str:
"""Custom renderer for Markdown paragraphs with MonsterUI styling."""
tokens[idx].attrSet("class", "leading-7 [&:not(:first-child)]:mt-6")
@@ -71,28 +72,30 @@ def render_paragraph(
def render_blockquote(
self: RendererHTML, tokens: List[Token], idx: int, options: dict, env: dict
self: RendererHTML, tokens: List[Token], idx: int, options: OptionsDict, env: dict
) -> str:
"""Custom renderer for Markdown blockquotes with MonsterUI styling."""
tokens[idx].attrSet("class", "mt-6 border-l-2 pl-6 italic border-primary")
return self.renderToken(tokens, idx, options, env)
def render_list(self: RendererHTML, tokens: List[Token], idx: int, options: dict, env: dict) -> str:
def render_list(
self: RendererHTML, tokens: List[Token], idx: int, options: OptionsDict, env: dict
) -> str:
"""Custom renderer for lists with MonsterUI styling."""
tokens[idx].attrSet("class", "my-6 ml-6 list-disc [&>li]:mt-2")
return self.renderToken(tokens, idx, options, env)
def render_image(
self: RendererHTML, tokens: List[Token], idx: int, options: dict, env: dict
self: RendererHTML, tokens: List[Token], idx: int, options: OptionsDict, env: dict
) -> str:
"""Custom renderer for Markdown images with MonsterUI styling."""
token = tokens[idx]
src = token.attrGet("src")
alt = token.content or ""
if src:
if isinstance(src, str) and src:
pos = src.find(ASSETS_PREFIX)
if pos != -1:
asset_rel = src[pos + len(ASSETS_PREFIX) :]
@@ -107,12 +110,14 @@ def render_image(
return self.renderToken(tokens, idx, options, env)
def render_link(self: RendererHTML, tokens: List[Token], idx: int, options: dict, env: dict) -> str:
def render_link(
self: RendererHTML, tokens: List[Token], idx: int, options: OptionsDict, env: dict
) -> str:
"""Custom renderer for Markdown links with MonsterUI styling."""
token = tokens[idx]
href = token.attrGet("href")
if href:
if isinstance(href, str) and href:
pos = href.find(ASSETS_PREFIX)
if pos != -1:
asset_rel = href[pos + len(ASSETS_PREFIX) :]
+12 -9
View File
@@ -3,7 +3,6 @@ from typing import Optional, Union
import pandas as pd
import requests
from bokeh.models import ColumnDataSource, LinearAxis, Range1d
from bokeh.plotting import figure
from loguru import logger
from monsterui.franken import (
Card,
@@ -28,7 +27,11 @@ from akkudoktoreos.core.emplan import (
OMBCInstruction,
)
from akkudoktoreos.optimization.optimization import OptimizationSolution
from akkudoktoreos.server.dash.bokeh import Bokeh, bokey_apply_theme_to_plot
from akkudoktoreos.server.dash.bokeh import (
Bokeh,
bokey_apply_theme_to_plot,
create_figure,
)
from akkudoktoreos.server.dash.components import Error
from akkudoktoreos.server.dash.context import request_url_for
from akkudoktoreos.utils.datetimeutil import compare_datetimes, to_datetime
@@ -259,21 +262,21 @@ def SolutionCard(solution: OptimizationSolution, config: SettingsEOS, data: Opti
last_run_datetime = "unknown"
start_datetime = "unknown"
plot = figure(
plot = create_figure(
title=f"Optimization Solution - last run: {last_run_datetime}",
x_axis_type="datetime",
x_axis_label=f"Datetime [localtime {date_time_tz}] - start: {start_datetime}",
y_axis_label="Power [W]",
sizing_mode="stretch_width",
y_range=Range1d(power_w_min, power_w_max),
y_range=Range1d(start=power_w_min, end=power_w_max),
height=400,
)
plot.extra_y_ranges = {
"energy": Range1d(energy_wh_min, energy_wh_max), # y2
"factor": Range1d(factor_min, factor_max), # y3
"amt_kwh": Range1d(amt_kwh_min, amt_kwh_max), # y4
"amt": Range1d(amt_min, amt_max), # y5
"energy": Range1d(start=energy_wh_min, end=energy_wh_max), # y2
"factor": Range1d(start=factor_min, end=factor_max), # y3
"amt_kwh": Range1d(start=amt_kwh_min, end=amt_kwh_max), # y4
"amt": Range1d(start=amt_min, end=amt_max), # y5
}
# y2 axis
y2_axis = LinearAxis(y_range_name="energy", axis_label="Energy [Wh]")
@@ -536,7 +539,7 @@ def InstructionCard(
)
):
# This is a battery
if instruction.operation_mode_id in ("CHARGE",):
if getattr(instruction, "operation_mode_id", None) in ("CHARGE",):
icon = "battery-charging"
else:
icon = "battery"
+10 -7
View File
@@ -3,11 +3,14 @@ from typing import Optional, Union
import pandas as pd
import requests
from bokeh.models import ColumnDataSource, LinearAxis, Range1d
from bokeh.plotting import figure
from monsterui.franken import FT, Grid, P
from akkudoktoreos.core.pydantic import PydanticDateTimeSeries
from akkudoktoreos.server.dash.bokeh import Bokeh, bokey_apply_theme_to_plot
from akkudoktoreos.server.dash.bokeh import (
Bokeh,
bokey_apply_theme_to_plot,
create_figure,
)
from akkudoktoreos.server.dash.components import Error
# bar width for 15 minutes bars (time given in millseconds)
@@ -18,7 +21,7 @@ def PVForecast(predictions: pd.DataFrame, config: dict, date_time_tz: str, dark:
source = ColumnDataSource(predictions)
provider = config["pvforecast"]["provider"]
plot = figure(
plot = create_figure(
x_axis_type="datetime",
title=f"PV Power Prediction ({provider})",
x_axis_label=f"Datetime [localtime {date_time_tz}]",
@@ -46,7 +49,7 @@ def ElectricityPriceForecast(
source = ColumnDataSource(predictions)
provider = config["elecprice"]["provider"]
plot = figure(
plot = create_figure(
x_axis_type="datetime",
y_range=Range1d(
predictions["elecprice_marketprice_kwh"].min() - 0.1,
@@ -78,7 +81,7 @@ def WeatherTempAirHumidityForecast(
source = ColumnDataSource(predictions)
provider = config["weather"]["provider"]
plot = figure(
plot = create_figure(
x_axis_type="datetime",
title=f"Air Temperature and Humidity Prediction ({provider})",
x_axis_label=f"Datetime [localtime {date_time_tz}]",
@@ -115,7 +118,7 @@ def WeatherIrradianceForecast(
source = ColumnDataSource(predictions)
provider = config["weather"]["provider"]
plot = figure(
plot = create_figure(
x_axis_type="datetime",
title=f"Irradiance Prediction ({provider})",
x_axis_label=f"Datetime [localtime {date_time_tz}]",
@@ -157,7 +160,7 @@ def LoadForecast(predictions: pd.DataFrame, config: dict, date_time_tz: str, dar
year_energy = config["load"]["loadakkudoktor"]["loadakkudoktor_year_energy_kwh"]
provider = f"{provider}, {year_energy} kWh"
plot = figure(
plot = create_figure(
title=f"Load Prediction ({provider})",
x_axis_type="datetime",
x_axis_label=f"Datetime [localtime {date_time_tz}]",
+58 -36
View File
@@ -35,6 +35,7 @@ from akkudoktoreos.core.coreabc import (
get_resource_registry,
singletons_init,
)
from akkudoktoreos.core.dataabc import DataImportMixin
from akkudoktoreos.core.emplan import EnergyManagementPlan, ResourceStatus
from akkudoktoreos.core.ems import ems_manage_energy
from akkudoktoreos.core.emsettings import EnergyManagementMode
@@ -88,7 +89,12 @@ from akkudoktoreos.server.server import (
get_host_ip,
wait_for_port_free,
)
from akkudoktoreos.utils.datetimeutil import to_datetime, to_duration
from akkudoktoreos.utils.datetimeutil import (
DateTime,
Duration,
to_datetime,
to_duration,
)
# ----------------------
# EOS REST Server
@@ -671,6 +677,8 @@ async def fastapi_logging_get_log(
"""
log_path = get_config().logging.file_path
try:
if log_path is None:
raise ValueError("Log file path is not configured")
logs = read_file_log(
log_path=log_path,
limit=limit,
@@ -864,16 +872,16 @@ async def fastapi_measurement_series_get(
if processing == SeriesProcessing.RAW:
pdseries = await get_measurement().key_to_raw_series(
key=key,
start_datetime=start_datetime,
end_datetime=end_datetime,
start_datetime=to_datetime(start_datetime) if start_datetime is not None else None,
end_datetime=to_datetime(end_datetime) if end_datetime is not None else None,
dropna=dropna,
)
else:
pdseries = await get_measurement().key_to_series(
key=key,
start_datetime=start_datetime,
end_datetime=end_datetime,
interval=interval,
start_datetime=to_datetime(start_datetime) if start_datetime is not None else None,
end_datetime=to_datetime(end_datetime) if end_datetime is not None else None,
interval=to_duration(interval) if interval is not None else None,
fill_method=fill_method,
resample_method=resample_method,
dropna=dropna,
@@ -1247,6 +1255,9 @@ async def fastapi_prediction_series_get(
Returns:
Array
"""
resolved_end_datetime: DateTime | None
resolved_interval: Duration
resolved_start_datetime: DateTime | None
if key not in get_prediction().record_keys:
raise EOSProblem(
status=404,
@@ -1255,10 +1266,10 @@ async def fastapi_prediction_series_get(
)
if start_datetime is None:
start_datetime = get_prediction().ems_start_datetime
resolved_start_datetime = get_prediction().ems_start_datetime
else:
try:
start_datetime = to_datetime(start_datetime)
resolved_start_datetime = to_datetime(start_datetime)
except Exception as e:
raise EOSProblem(
status=400,
@@ -1268,10 +1279,10 @@ async def fastapi_prediction_series_get(
) from e
if end_datetime is None:
end_datetime = get_prediction().end_datetime
resolved_end_datetime = get_prediction().end_datetime
else:
try:
end_datetime = to_datetime(end_datetime)
resolved_end_datetime = to_datetime(end_datetime)
except Exception as e:
raise EOSProblem(
status=400,
@@ -1281,10 +1292,10 @@ async def fastapi_prediction_series_get(
) from e
if interval is None:
interval = to_duration("1 hour")
resolved_interval = to_duration("1 hour")
else:
try:
interval = to_duration(interval)
resolved_interval = to_duration(interval)
except Exception as e:
raise EOSProblem(
status=400,
@@ -1297,16 +1308,16 @@ async def fastapi_prediction_series_get(
if processing == SeriesProcessing.RAW:
pdseries = await get_prediction().key_to_raw_series(
key=key,
start_datetime=start_datetime,
end_datetime=end_datetime,
start_datetime=resolved_start_datetime,
end_datetime=resolved_end_datetime,
dropna=dropna,
)
else:
pdseries = await get_prediction().key_to_series(
key=key,
start_datetime=start_datetime,
end_datetime=end_datetime,
interval=interval,
start_datetime=resolved_start_datetime,
end_datetime=resolved_end_datetime,
interval=resolved_interval,
fill_method=fill_method,
resample_method=resample_method,
dropna=dropna,
@@ -1417,24 +1428,26 @@ async def fastapi_prediction_dataframe_get(
forecast or reporting queries where alignment to the exact query window is
more important than clock-round boundaries.
"""
resolved_end_datetime: DateTime | None
resolved_start_datetime: DateTime | None
for key in keys:
if key not in get_prediction().record_keys:
raise HTTPException(status_code=404, detail=f"Key '{key}' is not available.")
if start_datetime is None:
start_datetime = get_prediction().ems_start_datetime
resolved_start_datetime = get_prediction().ems_start_datetime
else:
start_datetime = to_datetime(start_datetime)
resolved_start_datetime = to_datetime(start_datetime)
if end_datetime is None:
end_datetime = get_prediction().end_datetime
resolved_end_datetime = get_prediction().end_datetime
else:
end_datetime = to_datetime(end_datetime)
resolved_end_datetime = to_datetime(end_datetime)
try:
prediction_df = await get_prediction().keys_to_dataframe(
keys=keys,
start_datetime=start_datetime,
end_datetime=end_datetime,
interval=interval,
start_datetime=resolved_start_datetime,
end_datetime=resolved_end_datetime,
interval=to_duration(interval) if interval is not None else None,
fill_method=fill_method,
resample_method=resample_method,
dropna=dropna,
@@ -1536,6 +1549,9 @@ async def fastapi_prediction_list_get(
forecast or reporting queries where alignment to the exact query window is
more important than clock-round boundaries.
"""
resolved_end_datetime: DateTime | None
resolved_interval: Duration
resolved_start_datetime: DateTime | None
if key not in get_prediction().record_keys:
raise EOSProblem(
status=404,
@@ -1544,10 +1560,10 @@ async def fastapi_prediction_list_get(
)
if start_datetime is None:
start_datetime = get_prediction().ems_start_datetime
resolved_start_datetime = get_prediction().ems_start_datetime
else:
try:
start_datetime = to_datetime(start_datetime)
resolved_start_datetime = to_datetime(start_datetime)
except Exception as e:
raise EOSProblem(
status=400,
@@ -1557,10 +1573,10 @@ async def fastapi_prediction_list_get(
) from e
if end_datetime is None:
end_datetime = get_prediction().end_datetime
resolved_end_datetime = get_prediction().end_datetime
else:
try:
end_datetime = to_datetime(end_datetime)
resolved_end_datetime = to_datetime(end_datetime)
except Exception as e:
raise EOSProblem(
status=400,
@@ -1570,10 +1586,10 @@ async def fastapi_prediction_list_get(
) from e
if interval is None:
interval = to_duration("1 hour")
resolved_interval = to_duration("1 hour")
else:
try:
interval = to_duration(interval)
resolved_interval = to_duration(interval)
except Exception as e:
raise EOSProblem(
status=400,
@@ -1585,9 +1601,9 @@ async def fastapi_prediction_list_get(
try:
prediction_array = await get_prediction().key_to_array(
key=key,
start_datetime=start_datetime,
end_datetime=end_datetime,
interval=interval,
start_datetime=resolved_start_datetime,
end_datetime=resolved_end_datetime,
interval=resolved_interval,
fill_method=fill_method,
resample_method=resample_method,
dropna=dropna,
@@ -1655,6 +1671,12 @@ async def fastapi_prediction_import_provider(
cause=e,
) from e
if not isinstance(provider, DataImportMixin):
raise EOSProblem(
status=400,
title="Prediction import failed",
detail=f"Provider '{provider_id}' does not support data imports.",
)
await provider.import_from_json(json_str=json_str)
provider.update_datetime = to_datetime(in_timezone=get_config().general.timezone)
@@ -2192,8 +2214,8 @@ async def fastapi_optimize(
)
# Create compatible solution.
legacy_solution = Genetic0SolutionLegacy(
**{
legacy_solution = Genetic0SolutionLegacy.model_validate(
{
"ac_charge": solution.ac_charge,
"dc_charge": solution.dc_charge,
"discharge_allowed": solution.discharge_allowed,
@@ -2396,7 +2418,7 @@ def run_eos() -> None:
port=config_eos.server.port,
log_level=uv_log_level,
access_log=True, # Fix server access logging to True
reload=config_eos.server.reload,
reload=bool(config_eos.server.reload),
proxy_headers=True,
forwarded_allow_ips="*",
)
+12 -5
View File
@@ -1,6 +1,7 @@
import html
import traceback
from dataclasses import dataclass
from typing import cast
from fastapi import FastAPI, Request
from fastapi.exceptions import HTTPException, RequestValidationError
@@ -54,7 +55,9 @@ def _problem_response(
)
async def eos_problem_handler(request: Request, exc: EOSProblem) -> JSONResponse:
async def eos_problem_handler(request: Request, exc: Exception) -> JSONResponse:
# Starlette dispatches this handler by the registered exception class.
exc = cast(EOSProblem, exc)
return _problem_response(
request=request,
status=exc.status,
@@ -65,12 +68,14 @@ async def eos_problem_handler(request: Request, exc: EOSProblem) -> JSONResponse
)
async def http_exception_handler(request: Request, exc: HTTPException) -> JSONResponse:
async def http_exception_handler(request: Request, exc: Exception) -> JSONResponse:
# Starlette dispatches this handler by the registered exception class.
http_exc = cast(HTTPException, exc)
return _problem_response(
request=request,
status=exc.status_code,
status=http_exc.status_code,
title="HTTP Error",
detail=str(exc.detail),
detail=str(http_exc.detail),
cause=exc,
type="about:blank",
)
@@ -87,7 +92,9 @@ async def unexpected_exception_handler(request: Request, exc: Exception) -> JSON
)
async def validation_handler(request: Request, exc: RequestValidationError) -> JSONResponse:
async def validation_handler(request: Request, exc: Exception) -> JSONResponse:
# Starlette dispatches this handler by the registered exception class.
exc = cast(RequestValidationError, exc)
return _problem_response(
request=request,
status=422,
@@ -4,10 +4,14 @@ import re
import sys
import time
from pathlib import Path
from typing import Any, MutableMapping, Optional
from typing import TYPE_CHECKING, Any, MutableMapping, Optional
from loguru import logger
if TYPE_CHECKING:
from loguru import Record
from akkudoktoreos.core.coreabc import get_config
from akkudoktoreos.server.server import (
validate_ip_or_hostname,
@@ -99,7 +103,7 @@ def _emit_drop_warning() -> None:
def patch_loguru_record(
record: MutableMapping[str, Any],
record: "Record | MutableMapping[str, Any]",
*,
file_name: str,
file_path: str,
+6 -1
View File
@@ -124,7 +124,12 @@ def wait_for_port_free(port: int, timeout: int = 0, waiting_app_name: str = "App
try:
for conn in psutil.net_connections(kind="inet"):
if conn.laddr.port == port and conn.pid not in seen_pids:
if (
conn.laddr
and conn.laddr.port == port
and conn.pid is not None
and conn.pid not in seen_pids
):
try:
process = psutil.Process(conn.pid)
seen_pids.add(conn.pid)
+66 -31
View File
@@ -46,22 +46,36 @@ See each function's docstring for detailed argument options and examples.
import datetime
import re
from typing import Any, List, Literal, Optional, Tuple, Union, overload
from typing import (
TYPE_CHECKING,
Any,
Callable,
List,
Literal,
Optional,
Tuple,
Union,
cast,
overload,
)
import pendulum
from loguru import logger
from pendulum import UTC as UTC
from pendulum.tz.timezone import Timezone
from pydantic import (
GetCoreSchemaHandler,
)
from pydantic_core import core_schema
from pydantic_extra_types.pendulum_dt import ( # make pendulum types pydantic
Date,
DateTime,
Duration,
)
from tzfpy import get_tz
if TYPE_CHECKING:
# The Pydantic adapters validate Pendulum values; arithmetic and factory
# functions return the base types rather than the validation subclasses.
from pendulum import Date, DateTime, Duration
else:
from pydantic_extra_types.pendulum_dt import Date, DateTime, Duration
MAX_DURATION_STRING_LENGTH = 350
@@ -184,7 +198,7 @@ class Time(pendulum.Time):
# Bypass __init__ and __new__ by directly casting the type
time_obj.__class__ = cls # This is safe since Time inherits from pendulum.Time
return time_obj
return cast(Time, time_obj)
@classmethod
def _serialize(cls, value: Optional["Time"]) -> str:
@@ -224,7 +238,7 @@ class Time(pendulum.Time):
if self.tzinfo and other.tzinfo:
# Convert both to UTC for comparison
self_utc = self.in_timezone("UTC")
other_utc = other.in_timezone("UTC")
other_utc = cast(Time, other).in_timezone("UTC")
return (self_utc.hour, self_utc.minute, self_utc.second, self_utc.microsecond) == (
other_utc.hour,
other_utc.minute,
@@ -259,7 +273,9 @@ class Time(pendulum.Time):
"""Convert to UTC timezone."""
return self.in_timezone("UTC")
def in_timezone(self, timezone: Union[str, pendulum.Timezone]) -> "Time":
def in_timezone(
self, timezone: Union[str, pendulum.Timezone, pendulum.FixedTimezone]
) -> "Time":
"""Convert to specified timezone."""
if isinstance(timezone, str):
timezone = pendulum.timezone(timezone)
@@ -267,7 +283,9 @@ class Time(pendulum.Time):
if self.is_aware():
# For timezone conversion, we need a reference date
# Use today's date as reference
today = pendulum.today(self.tzinfo)
today = cast(Callable[[datetime.tzinfo | None], pendulum.DateTime], pendulum.today)(
self.tzinfo
)
dt = today.at(self.hour, self.minute, self.second, self.microsecond)
dt = dt.in_timezone(timezone) # Convert to target timezone
t = dt.time() # Extract naiv time component
@@ -316,7 +334,7 @@ class Time(pendulum.Time):
return self.format(time_format)
@classmethod
def now(cls, tz: Union[str, pendulum.Timezone] = None) -> "Time":
def now(cls, tz: Union[str, pendulum.Timezone, None] = None) -> "Time":
"""Get current time with optional timezone."""
if tz:
if isinstance(tz, str):
@@ -336,7 +354,7 @@ class Time(pendulum.Time):
)
def _parse_time_string(time_str: str, default_date: pendulum.Date = None) -> pendulum.Time:
def _parse_time_string(time_str: str, default_date: pendulum.Date | None = None) -> pendulum.Time:
"""Parse various time string formats with comprehensive patterns and timezone support.
Supports a wide variety of time formats including:
@@ -387,7 +405,7 @@ def _parse_time_string(time_str: str, default_date: pendulum.Date = None) -> pen
raise ValueError("Empty time string")
# Extract timezone information first
timezone_info = None
timezone_info: pendulum.Timezone | pendulum.FixedTimezone | None = None
time_part = time_str
# Pattern for timezone at the end: +HH:MM, -HH:MM, +HHMM, -HHMM, UTC, GMT, EST, PST, etc.
@@ -703,7 +721,9 @@ def to_time(
# Convert from original timezone to selected timezone
# For timezone conversion, we need a reference date
# Use today's date as reference
today = pendulum.today(t.tzinfo)
today = cast(
Callable[[datetime.tzinfo | None], pendulum.DateTime], pendulum.today
)(t.tzinfo)
dt = today.at(t.hour, t.minute, t.second, t.microsecond)
dt = dt.in_timezone(timezone) # Convert to target timezone
t = dt.time() # Extract time component (always naive)
@@ -746,7 +766,7 @@ def to_time(
tz_name = value.tzinfo.tzname(value)
# Safely get Pendulum timezone
try:
timezone = pendulum.timezone(tz_name)
timezone = pendulum.timezone(cast(str, tz_name))
except Exception:
# fallback to fixed offset if tz_name is something like 'UTC+02:00'
utc_offset = value.tzinfo.utcoffset(value)
@@ -754,7 +774,7 @@ def to_time(
utc_offset_total_seconds = 0.0
else:
utc_offset_total_seconds = utc_offset.total_seconds()
timezone = pendulum.FixedTimezone(utc_offset_total_seconds // 60)
timezone = pendulum.FixedTimezone(int(utc_offset_total_seconds // 60))
pdt = pendulum.instance(value).in_tz(timezone)
return finalize(pdt.time())
@@ -792,21 +812,25 @@ def to_time(
# Fallback to pendulum's parser
try:
dt = pendulum.parse(value, strict=False).in_tz(timezone)
dt = cast(pendulum.DateTime, pendulum.parse(value, strict=False)).in_tz(timezone)
return finalize(dt.time())
except Exception as e:
logger.trace(f"Pendulum parser failed for '{value}': {e}")
# Try parsing with ISO time prefix
try:
dt = pendulum.parse(f"T{value}", strict=False).in_tz(timezone)
dt = cast(pendulum.DateTime, pendulum.parse(f"T{value}", strict=False)).in_tz(
timezone
)
return finalize(dt.time())
except Exception as e:
logger.trace(f"ISO time parser failed for 'T{value}': {e}")
# Try parsing as part of a full datetime
try:
dt = pendulum.parse(f"2000-01-01 {value}", strict=False).in_tz(timezone)
dt = cast(
pendulum.DateTime, pendulum.parse(f"2000-01-01 {value}", strict=False)
).in_tz(timezone)
return finalize(dt.time())
except Exception as e:
logger.trace(f"Full datetime parser failed for '2000-01-01 {value}': {e}")
@@ -903,16 +927,19 @@ def to_datetime(
'2024-10-31 12:00:00'
"""
# Timezone to convert to
timezone: Timezone | pendulum.FixedTimezone
if in_timezone is None:
in_timezone = pendulum.local_timezone()
elif not isinstance(in_timezone, Timezone):
in_timezone = pendulum.timezone(in_timezone)
timezone = pendulum.local_timezone()
elif isinstance(in_timezone, Timezone):
timezone = in_timezone
else:
timezone = pendulum.timezone(in_timezone)
if isinstance(date_input, DateTime):
dt = date_input
elif isinstance(date_input, Date):
dt = pendulum.datetime(
year=date_input.year, month=date_input.month, day=date_input.day, tz=in_timezone
year=date_input.year, month=date_input.month, day=date_input.day, tz=timezone
)
if to_maxtime:
dt = dt.end_of("day")
@@ -937,10 +964,10 @@ def to_datetime(
# DateTime input without timezone info
try:
fmt_tz = f"{fmt} z"
dt_tz = f"{date_input} {in_timezone}"
dt_tz = f"{date_input} {timezone}"
dt = pendulum.from_format(dt_tz, fmt_tz)
logger.trace(
f"Str Fmt converted: {dt}, tz={dt.tz} from {date_input}, tz={in_timezone}"
f"Str Fmt converted: {dt}, tz={dt.tz} from {date_input}, tz={timezone}"
)
break
except ValueError as e:
@@ -949,9 +976,9 @@ def to_datetime(
else:
# DateTime input with timezone info
try:
dt = pendulum.parse(date_input)
dt = cast(pendulum.DateTime, pendulum.parse(date_input))
logger.trace(
f"Pendulum Fmt converted: {dt}, tz={dt.tz} from {date_input}, tz={in_timezone}"
f"Pendulum Fmt converted: {dt}, tz={dt.tz} from {date_input}, tz={timezone}"
)
except pendulum.parsing.exceptions.ParserError as e:
logger.trace(f"Date string {date_input} does not match any Pendulum formats: {e}")
@@ -971,7 +998,9 @@ def to_datetime(
if dt is None:
raise ValueError(f"Date string {date_input} does not match any known formats.")
elif date_input is None:
dt = pendulum.now(tz=in_timezone)
dt = cast(Callable[[Timezone | pendulum.FixedTimezone], pendulum.DateTime], pendulum.now)(
timezone
)
elif isinstance(date_input, datetime.datetime):
dt = pendulum.instance(date_input)
elif isinstance(date_input, datetime.date):
@@ -988,10 +1017,14 @@ def to_datetime(
logger.error(error_msg)
raise ValueError(error_msg)
# Every supported input branch produces a datetime or raises above.
if dt is None:
raise ValueError("Datetime conversion did not produce a value")
# Represent in target timezone
dt_in_tz = dt.in_timezone(in_timezone)
dt_in_tz = dt.in_timezone(timezone)
logger.trace(
f"\nTimezone adapted to: {in_timezone}\nfrom: {dt} tz={dt.timezone}\nto: {dt_in_tz} tz={dt_in_tz.tz}"
f"\nTimezone adapted to: {timezone}\nfrom: {dt} tz={dt.timezone}\nto: {dt_in_tz} tz={dt_in_tz.tz}"
)
dt = dt_in_tz
@@ -1158,7 +1191,8 @@ def to_duration(
duration = parsed # Already a duration
else:
# It's a DateTime, calculate duration from start of day
duration = parsed - parsed.start_of("day")
parsed_datetime = cast(pendulum.DateTime, parsed)
duration = parsed_datetime - parsed_datetime.start_of("day")
except pendulum.parsing.exceptions.ParserError as e:
logger.trace(f"Invalid Pendulum time string format '{input_value}': {e}")
@@ -1516,6 +1550,7 @@ def compare_datetimes(
DatetimesComparisonResult(equal=False, same_instant=True, time_diff=7200, timezone_diff=True, dst_diff=False, approximately_equal=True, ge=False, gt=False, le=True, lt=True)
"""
# Normalize tolerance to seconds
tolerance_seconds: float
if tolerance is None:
tolerance_seconds = 0
elif isinstance(tolerance, pendulum.Duration):