import unittest from unittest.mock import patch import database class FakeResponse: def __init__(self, data=None): self.data = data or [] class FakeRpc: def __init__(self, client, params): self.client = client self.params = params def execute(self): if self.client.fail_rpc: raise RuntimeError("rpc failed") saved = [] for row in self.params["p_records"]: key = (row["plant_id"], row["date"]) existing = self.client.daily.get(key) incoming = row["total_generation"] if ( existing is None or self.params["p_allow_decrease"] or incoming > existing ): self.client.daily[key] = incoming saved.append({**row, "total_generation": self.client.daily[key]}) self.client.rpc_calls.append(self.params) return FakeResponse(saved) class FakeQuery: def __init__(self, client, table_name): self.client = client self.table_name = table_name self.action = None self.payload = None self.filters = {} def select(self, _columns): self.action = "select" return self def upsert(self, payload, on_conflict=None): self.action = "upsert" self.payload = payload self.on_conflict = on_conflict return self def eq(self, column, value): self.filters[column] = ("eq", value) return self def in_(self, column, values): self.filters[column] = ("in", set(values)) return self def gte(self, column, value): self.filters[f"{column}_gte"] = ("gte", value) return self def lte(self, column, value): self.filters[f"{column}_lte"] = ("lte", value) return self def execute(self): if self.client.fail_select and self.action == "select": raise RuntimeError("select failed") if self.action == "select" and self.table_name == "daily_stats": rows = [] for (plant_id, date), value in self.client.daily.items(): if "plant_id" in self.filters and plant_id != self.filters["plant_id"][1]: continue if "date" in self.filters and date not in self.filters["date"][1]: continue if "date_gte" in self.filters and date < self.filters["date_gte"][1]: continue if "date_lte" in self.filters and date > self.filters["date_lte"][1]: continue rows.append({ "plant_id": plant_id, "date": date, "total_generation": value, }) return FakeResponse(rows) if self.action == "upsert" and self.table_name == "daily_stats": for row in self.payload: self.client.daily[(row["plant_id"], row["date"])] = row["total_generation"] self.client.daily_upserts.append(self.payload) return FakeResponse(self.payload) if self.action == "upsert" and self.table_name == "monthly_stats": self.client.monthly_upserts.append(self.payload) return FakeResponse(self.payload) raise AssertionError(f"unexpected query: {self.table_name} {self.action}") class FakeClient: def __init__(self, daily=None, fail_rpc=False): self.daily = daily or {} self.fail_select = False self.fail_rpc = fail_rpc self.daily_upserts = [] self.monthly_upserts = [] self.rpc_calls = [] def table(self, table_name): return FakeQuery(self, table_name) def rpc(self, function_name, params): if function_name != "upsert_daily_stats": raise AssertionError(function_name) return FakeRpc(self, params) class DatabaseHistoryTest(unittest.TestCase): def save_daily(self, client, generation): with patch.object(database, "get_supabase_client", return_value=client): return database.save_history([{ "plant_id": "plant-a", "date": "2026-08-05", "generation_kwh": generation, }], "daily") def test_existing_positive_value_is_not_replaced_by_zero(self): client = FakeClient({("plant-a", "2026-08-05"): 100.0}) result = self.save_daily(client, 0.0) self.assertTrue(result) self.assertEqual(client.daily[("plant-a", "2026-08-05")], 100.0) self.assertEqual(client.daily_upserts, []) self.assertEqual(client.rpc_calls[0]["p_source"], "history") self.assertFalse(client.rpc_calls[0]["p_allow_decrease"]) def test_larger_value_updates_existing_daily_stat(self): client = FakeClient({("plant-a", "2026-08-05"): 100.0}) result = self.save_daily(client, 120.0) self.assertTrue(result) self.assertEqual(client.daily[("plant-a", "2026-08-05")], 120.0) self.assertEqual(len(client.rpc_calls), 1) def test_rpc_failure_prevents_write(self): client = FakeClient(fail_rpc=True) result = self.save_daily(client, 120.0) self.assertFalse(result) self.assertEqual(client.rpc_calls, []) def test_negative_generation_is_rejected(self): client = FakeClient() result = self.save_daily(client, -1.0) self.assertFalse(result) self.assertEqual(client.rpc_calls, []) if __name__ == "__main__": unittest.main()