mirror of
https://github.com/Akkudoktor-EOS/EOS.git
synced 2026-10-08 15:26:38 +00:00
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:
co-authored by
dr-dimitri
Normann
parent
5b584cbb57
commit
1abdd345c4
@@ -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[
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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."},
|
||||
)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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))
|
||||
]
|
||||
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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],
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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"},
|
||||
)
|
||||
|
||||
@@ -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"}
|
||||
)
|
||||
|
||||
@@ -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]]:
|
||||
|
||||
@@ -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",
|
||||
)
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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) :]
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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}]",
|
||||
|
||||
@@ -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="*",
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user