diff --git a/README.md b/README.md index c3667ee..f393586 100644 --- a/README.md +++ b/README.md @@ -617,6 +617,13 @@ Forecaster represents a generative source of predicted values. | Field | Type | Label | Description | | ----- | ---- | ----- | ----------- | | observer_uuid | [string](#string) | | || observer_name | [string](#string) | | | + +
StreamCreateForecastsResponse + + +| Field | Type | Label | Description | +| ----- | ---- | ----- | ----------- | +| forecast_uuids | [string](#string) | repeated | A list of the UUIDs generated for the successfully created forecasts. |
StreamForecastDataRequest @@ -886,6 +893,15 @@ Useful for analytics and performance monitoring. _[StreamForecastDataRequest](#ocf-dp-StreamForecastDataRequest) / [StreamForecastDataResponse](#ocf-dp-StreamForecastDataResponse) stream_ + + +#### StreamCreateForecasts + +StreamCreateForecasts allows for efficient batch creation of multiple forecasts and their values. +Note: This method is executed in a single transaction. To prevent resource exhaustion, a maximum of 5000 forecasts can be sent per stream. Exceeding this limit will abort the stream and roll back all inserts. + +_[CreateForecastRequest](#ocf-dp-CreateForecastRequest) stream / [StreamCreateForecastsResponse](#ocf-dp-StreamCreateForecastsResponse)_ + diff --git a/examples/python-notebook/example.py b/examples/python-notebook/example.py index e00721b..24b26f5 100644 --- a/examples/python-notebook/example.py +++ b/examples/python-notebook/example.py @@ -20,7 +20,6 @@ from ocf.dp.dp import common_pb2 from ocf.dp.dp_data import messages_pb2, service_pb2_grpc import pandas as pd -import xarray as xr import datetime as dt @@ -77,7 +76,7 @@ async def main() -> None: end_time = gfreq_response.values[-1].target_timestamp_utc.ToDatetime(tzinfo=dt.UTC) print(f"\tReceived {len(gfreq_response.values)} forecast points from {start_time} to {end_time}") - print(f":: -> Converting response to a dataframe") + print(":: -> Converting response to a dataframe") # preserving_proto_field_name prevents conversion to lowerCamelCase. # always_print_fields_with_no_presence ensures all fields are present in the dict, even if they have no value in the protobuf. df = pd.DataFrame.from_dict([ @@ -95,14 +94,14 @@ async def main() -> None: ).drop(["p50_value_fraction", "p10", "p90"], axis=1) print(df.head()) - print(f":: Getting 'ground truths' for the same location and time period") + print(":: Getting 'ground truths' for the same location and time period") - print(f":: -> Getting an observer") + print(":: -> Getting an observer") loresp = await dpc.ListObservers(messages_pb2.ListObserversRequest()) observer = next(o for o in loresp.observers if "pvlive" in o.observer_name) print(f"\t{observer.observer_name=}") - print(f":: -> Getting the ground truth for the UK national location") + print(":: -> Getting the ground truth for the UK national location") gtreq = messages_pb2.GetObservationsAsTimeseriesRequest( location_uuid=uk_location.location_uuid, energy_source=common_pb2.EnergySource.ENERGY_SOURCE_SOLAR, diff --git a/go.mod b/go.mod index ca69b64..4fa265b 100644 --- a/go.mod +++ b/go.mod @@ -13,6 +13,7 @@ require ( github.com/rs/zerolog v1.34.0 github.com/stretchr/testify v1.11.1 github.com/testcontainers/testcontainers-go v0.40.0 + golang.org/x/sync v0.20.0 google.golang.org/grpc v1.79.2 google.golang.org/protobuf v1.36.11 ) @@ -80,7 +81,6 @@ require ( golang.org/x/crypto v0.48.0 // indirect golang.org/x/exp v0.0.0-20260218203240-3dfff04db8fa // indirect golang.org/x/net v0.51.0 // indirect - golang.org/x/sync v0.20.0 // indirect golang.org/x/sys v0.42.0 // indirect golang.org/x/text v0.34.0 // indirect golang.org/x/time v0.12.0 // indirect diff --git a/internal/server/dummy/dataserverimpl.go b/internal/server/dummy/dataserverimpl.go index c93146e..e2bf5d7 100644 --- a/internal/server/dummy/dataserverimpl.go +++ b/internal/server/dummy/dataserverimpl.go @@ -18,6 +18,8 @@ import ( "github.com/google/uuid" "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" "google.golang.org/protobuf/types/known/structpb" "google.golang.org/protobuf/types/known/timestamppb" @@ -723,5 +725,12 @@ func (d *DataPlatformDataServiceServerImpl) UpdateForecaster( }, nil } +// StreamCreateForecasts implements dp.DataPlatformDataServiceServer. +func (s *DataPlatformDataServiceServerImpl) StreamCreateForecasts( + stream grpc.ClientStreamingServer[pb.CreateForecastRequest, pb.StreamCreateForecastsResponse], +) error { + return status.Errorf(codes.Unimplemented, "method StreamCreateForecasts not implemented") +} + // Compile-time check to ensure the interface is implemented fully. var _ pb.DataPlatformDataServiceServer = (*DataPlatformDataServiceServerImpl)(nil) diff --git a/internal/server/postgres/dataserverimpl.go b/internal/server/postgres/dataserverimpl.go index b303176..a4382fe 100644 --- a/internal/server/postgres/dataserverimpl.go +++ b/internal/server/postgres/dataserverimpl.go @@ -10,6 +10,7 @@ import ( "context" "errors" "fmt" + "io" "time" "github.com/google/uuid" @@ -75,6 +76,66 @@ func extractSIPStatPtrFromMap(m map[string]float32, key string) *int16 { return &sip_val } +// prepareForecastParams generates the database parameters for a single forecast from a gRPC request. +func prepareForecastParams( + req *pb.CreateForecastRequest, + geometryUuid uuid.UUID, + sourceTypeId int16, + forecasterId int32, +) (db.CreateForecastsParams, error) { + initTime := req.InitTimeUtc.AsTime().Truncate(time.Minute) + + fUuid, err := uuid.NewV7() + if err != nil { + return db.CreateForecastsParams{}, fmt.Errorf("failed to generate uuidv7: %w", err) + } + + // Manually overwrite the 48-bit timestamp with the initTime milliseconds + ms := uint64(initTime.UnixMilli()) + fUuid[0] = byte(ms >> 40) + fUuid[1] = byte(ms >> 32) + fUuid[2] = byte(ms >> 24) + fUuid[3] = byte(ms >> 16) + fUuid[4] = byte(ms >> 8) + fUuid[5] = byte(ms) + + firstHorizon := int32(req.Values[0].HorizonMins) + lastHorizon := int32(req.Values[len(req.Values)-1].HorizonMins) + + periodStart := initTime.Add(time.Duration(firstHorizon) * time.Minute) + periodEnd := initTime.Add(time.Duration(lastHorizon) * time.Minute) + + targetPeriod := pgtype.Range[pgtype.Timestamp]{ + Lower: pgtype.Timestamp{Time: periodStart, Valid: true}, + Upper: pgtype.Timestamp{Time: periodEnd, Valid: true}, + LowerType: pgtype.Inclusive, + UpperType: pgtype.Inclusive, + Valid: true, + } + + var createdTime pgtype.Timestamp + if req.CreatedTimestampUtc != nil { + createdTime = pgtype.Timestamp{Time: req.CreatedTimestampUtc.AsTime(), Valid: true} + } else { + createdTime = pgtype.Timestamp{ + Time: time.Now().UTC().Truncate(time.Minute), + Valid: true, + } + } + + return db.CreateForecastsParams{ + ForecastUuid: fUuid, + GeometryUuid: geometryUuid, + SourceTypeID: sourceTypeId, + ForecasterID: forecasterId, + InitTimeUtc: pgtype.Timestamp{Time: initTime, Valid: true}, + ValueResolutionMins: int16(req.Values[1].HorizonMins - req.Values[0].HorizonMins), + TargetPeriod: targetPeriod, + Metadata: req.Metadata, + CreatedAtUtc: createdTime, + }, nil +} + // --- Server Implementation ---------------------------------------------------------------------- func NewDataPlatformDataServiceServerImpl() *DataPlatformDataServiceServerImpl { @@ -146,21 +207,22 @@ func (s *DataPlatformDataServiceServerImpl) CreateForecast( Msg("found forecaster") // Create a new forecast - cfprms := db.CreateForecastParams{ - GeometryUuid: uuid.MustParse(req.LocationUuid), - SourceTypeID: dbSource.SourceTypeID, - ForecasterID: dbForecaster.ForecasterID, - ValueResolutionMins: int16(resolution_mins), - InitTimeUtc: timeptrToPgTimestamp(req.InitTimeUtc), - FirstHorizonMins: int32(req.Values[0].HorizonMins), - // Okay to take the last value as we checked it was monotonically increasing above - LastHorizonMins: int32(req.Values[len(req.Values)-1].HorizonMins), - Metadata: req.Metadata, - CreatedAtUtc: timeptrToPgTimestamp(req.CreatedTimestampUtc), - } - - dbForecast, err := querier.CreateForecast(ctx, cfprms) + fParams, err := prepareForecastParams( + req, + uuid.MustParse(req.LocationUuid), + dbSource.SourceTypeID, + dbForecaster.ForecasterID, + ) if err != nil { + return nil, fmt.Errorf("failed to prepare forecast params: %w", err) + } + + countF, err := querier.CreateForecasts(ctx, []db.CreateForecastsParams{fParams}) + if err != nil || countF < 1 { + if err == nil { + err = errors.New("inserted forecasts count less than requested") + } + return nil, fmt.Errorf("invalid forecast: %w", err) } @@ -176,7 +238,7 @@ func (s *DataPlatformDataServiceServerImpl) CreateForecast( P75Sip: extractSIPStatPtrFromMap(value.OtherStatisticsFractions, "p75"), P90Sip: extractSIPStatPtrFromMap(value.OtherStatisticsFractions, "p90"), P98Sip: extractSIPStatPtrFromMap(value.OtherStatisticsFractions, "p98"), - ForecastUuid: dbForecast.ForecastUuid, + ForecastUuid: fParams.ForecastUuid, } } @@ -190,17 +252,17 @@ func (s *DataPlatformDataServiceServerImpl) CreateForecast( } l.Debug(). - Str("dp.forecast.uuid", dbForecast.ForecastUuid.String()). - Str("dp.geometry.uuid", dbForecast.GeometryUuid.String()). - Str("dp.forecast.init_time", dbForecast.InitTimeUtc.Time.String()). + Str("dp.forecast.uuid", fParams.ForecastUuid.String()). + Str("dp.geometry.uuid", fParams.GeometryUuid.String()). + Str("dp.forecast.init_time", fParams.InitTimeUtc.Time.String()). Str("dp.forecast.target_period", fmt.Sprintf( "%s - %s", - dbForecast.TargetPeriod.Lower.Time.String(), - dbForecast.TargetPeriod.Upper.Time.String(), + fParams.TargetPeriod.Lower.Time.String(), + fParams.TargetPeriod.Upper.Time.String(), )).Msgf("created forecast") return &pb.CreateForecastResponse{ - ForecastUuid: dbForecast.ForecastUuid.String(), + ForecastUuid: fParams.ForecastUuid.String(), }, nil } @@ -1782,5 +1844,213 @@ func (s *DataPlatformDataServiceServerImpl) ListLocations( }, nil } +// StreamCreateForecasts efficiently creates multiple forecasts and their predictions via copyfrom batching. +func (s *DataPlatformDataServiceServerImpl) StreamCreateForecasts( + stream grpc.ClientStreamingServer[pb.CreateForecastRequest, pb.StreamCreateForecastsResponse], +) error { + ctx := stream.Context() + pool := ix.GetPoolFromContext(ctx) + + tx, err := pool.Begin(ctx) + querier := db.New(tx) + + if err != nil { + return status.Errorf(codes.Internal, "failed to begin transaction: %v", err) + } + + defer func() { _ = tx.Rollback(ctx) }() + + const ( + batchSize = 500 + maxBatches = 10 + ) + batchesProcessed := 0 + + var ( + forecastParams []db.CreateForecastsParams + valueParams []db.CreatePredictedValuesParams + createdUuids []string + batchUuids []string + ) + + // In-memory caches to avoid hammering the database for repeated forecaster/source lookups + type sourceKey struct { + locationUuid string + sourceTypeId int16 + } + + type sourceInfo struct { + capacityWatts int64 + geometryUuid uuid.UUID + } + + type forecasterKey struct { + name string + version string + } + + sourceCache := make(map[sourceKey]sourceInfo) + forecasterCache := make(map[forecasterKey]int32) + + flushBatch := func() error { + if len(forecastParams) == 0 { + return nil + } + + countF, err := querier.CreateForecasts(ctx, forecastParams) + if err != nil || countF < int64(len(forecastParams)) { + if err == nil { + err = errors.New("inserted forecasts count less than requested") + } + + return fmt.Errorf("failed to insert forecasts batch: %w", err) + } + + countV, err := querier.CreatePredictedValues(ctx, valueParams) + if err != nil || countV < int64(len(valueParams)) { + if err == nil { + err = errors.New("inserted predicted values count less than requested") + } + + return fmt.Errorf("failed to insert predicted values batch: %w", err) + } + + createdUuids = append(createdUuids, batchUuids...) + + // Reset batch buffers + forecastParams = forecastParams[:0] + valueParams = valueParams[:0] + batchUuids = batchUuids[:0] + + batchesProcessed++ + if batchesProcessed > maxBatches { + return status.Error( + codes.InvalidArgument, + fmt.Sprintf( + "maximum number of forecasts per stream exceeded (%d)", + maxBatches*batchSize, + ), + ) + } + + return nil + } + + for { + req, err := stream.Recv() + if err != nil { + if errors.Is(err, io.EOF) { + break + } + + return fmt.Errorf("error receiving from stream: %w", err) + } + + fKey := forecasterKey{ + name: req.Forecaster.ForecasterName, + version: req.Forecaster.ForecasterVersion, + } + + fId, ok := forecasterCache[fKey] + if !ok { + pctprms := db.GetForecasterElseLatestParams{ + ForecasterName: fKey.name, + ForecasterVersion: fKey.version, + } + + dbForecaster, err := querier.GetForecasterElseLatest(ctx, pctprms) + if err != nil { + return fmt.Errorf( + "no forecaster found for name '%s' and version '%s': %w", + fKey.name, + fKey.version, + err, + ) + } + + fId = dbForecaster.ForecasterID + forecasterCache[fKey] = fId + } + + sKey := sourceKey{ + locationUuid: req.LocationUuid, + sourceTypeId: int16(req.EnergySource.Number()), + } + + sInfo, ok := sourceCache[sKey] + if !ok { + gsprms := db.GetSourceAtTimestampParams{ + GeometryUuid: uuid.MustParse(req.LocationUuid), + SourceTypeID: sKey.sourceTypeId, + AtTimestampUtc: timeptrToPgTimestamp(req.InitTimeUtc), + } + + dbSource, err := querier.GetSourceAtTimestamp(ctx, gsprms) + if err != nil { + return fmt.Errorf( + "no location source found for name '%s' with source type '%s': %w", + req.LocationUuid, + req.EnergySource, + err, + ) + } + + sInfo = sourceInfo{ + capacityWatts: dbSource.CapacityWatts, + geometryUuid: dbSource.GeometryUuid, + } + sourceCache[sKey] = sInfo + } + + fParams, err := prepareForecastParams( + req, + sInfo.geometryUuid, + sKey.sourceTypeId, + fId, + ) + if err != nil { + return fmt.Errorf("failed to prepare forecast params: %w", err) + } + + forecastParams = append(forecastParams, fParams) + + for _, value := range req.Values { + valueParams = append(valueParams, db.CreatePredictedValuesParams{ + HorizonMins: int16(value.HorizonMins), + P02Sip: extractSIPStatPtrFromMap(value.OtherStatisticsFractions, "p02"), + P10Sip: extractSIPStatPtrFromMap(value.OtherStatisticsFractions, "p10"), + P25Sip: extractSIPStatPtrFromMap(value.OtherStatisticsFractions, "p25"), + P50Sip: int16(value.P50Fraction * 30000.0), + P75Sip: extractSIPStatPtrFromMap(value.OtherStatisticsFractions, "p75"), + P90Sip: extractSIPStatPtrFromMap(value.OtherStatisticsFractions, "p90"), + P98Sip: extractSIPStatPtrFromMap(value.OtherStatisticsFractions, "p98"), + ForecastUuid: fParams.ForecastUuid, + }) + } + + batchUuids = append(batchUuids, fParams.ForecastUuid.String()) + + // Flush if we hit the batch size limit + if len(forecastParams) >= batchSize { + if err := flushBatch(); err != nil { + return err + } + } + } + + // Flush any remaining requests + if err := flushBatch(); err != nil { + return err + } + + if err := tx.Commit(ctx); err != nil { + return status.Errorf(codes.Internal, "failed to commit transaction: %v", err) + } + + return stream.SendAndClose(&pb.StreamCreateForecastsResponse{ + ForecastUuids: createdUuids, + }) +} + // Compile-time check to ensure the interface is implemented fully. var _ pb.DataPlatformDataServiceServer = (*DataPlatformDataServiceServerImpl)(nil) diff --git a/internal/server/postgres/dataserverimpl_test.go b/internal/server/postgres/dataserverimpl_test.go index 3a9780e..1d4ddbc 100644 --- a/internal/server/postgres/dataserverimpl_test.go +++ b/internal/server/postgres/dataserverimpl_test.go @@ -1,6 +1,7 @@ package postgres import ( + "context" "encoding/hex" "encoding/json" "fmt" @@ -12,6 +13,8 @@ import ( "github.com/google/uuid" "github.com/stretchr/testify/require" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" "google.golang.org/protobuf/types/known/structpb" timestamppb "google.golang.org/protobuf/types/known/timestamppb" @@ -2439,3 +2442,301 @@ func TestStreamForecastData(t *testing.T) { }) } } + +func TestStreamCreateForecasts(t *testing.T) { + pivotTime := time.Date(2024, 6, 1, 0, 0, 0, 0, time.UTC) + + // Create a site + siteResp := createTestLocation( + t, + "test_stream_create_forecasts_site", + "POINT(-0.1 51.5)", + 1000000, + pivotTime.Add(-time.Hour*24), + nil, + ) + + // Create a forecaster + fc := createTestForecaster(t, "test_stream_create_forecasts_forecaster", "v1") + + yields := generateTestForecastValues(10, 30) + + testcases := []struct { + name string + setupStream func(ctx context.Context) (pb.DataPlatformDataService_StreamCreateForecastsClient, error) + sendCount int + getReq func(i int) *pb.CreateForecastRequest + shouldErr bool + expectedErrCode codes.Code + expectedUuidsCount int + }{ + { + name: "Valid stream under limit", + setupStream: func(ctx context.Context) (pb.DataPlatformDataService_StreamCreateForecastsClient, error) { + return dc.StreamCreateForecasts(ctx) + }, + sendCount: 10, + getReq: func(i int) *pb.CreateForecastRequest { + return &pb.CreateForecastRequest{ + LocationUuid: siteResp.LocationUuid, + Forecaster: &pb.Forecaster{ + ForecasterName: fc.ForecasterName, + ForecasterVersion: fc.ForecasterVersion, + }, + EnergySource: pb.EnergySource_ENERGY_SOURCE_SOLAR, + InitTimeUtc: timestamppb.New(pivotTime.Add(time.Duration(i) * time.Hour)), + Values: yields, + } + }, + shouldErr: false, + expectedUuidsCount: 10, + }, + { + name: "Atomicity on Failure (rollback after valid batch)", + setupStream: func(ctx context.Context) (pb.DataPlatformDataService_StreamCreateForecastsClient, error) { + return dc.StreamCreateForecasts(ctx) + }, + sendCount: 600, // Should exceed batch size of 500 + getReq: func(i int) *pb.CreateForecastRequest { + // Inject error at the end + if i == 599 { + return &pb.CreateForecastRequest{ + LocationUuid: siteResp.LocationUuid, + Forecaster: &pb.Forecaster{ + ForecasterName: "non_existent", + ForecasterVersion: "v1", + }, + EnergySource: pb.EnergySource_ENERGY_SOURCE_SOLAR, + InitTimeUtc: timestamppb.New(pivotTime), + Values: yields, + } + } + + return &pb.CreateForecastRequest{ + LocationUuid: siteResp.LocationUuid, + Forecaster: &pb.Forecaster{ + ForecasterName: fc.ForecasterName, + ForecasterVersion: fc.ForecasterVersion, + }, + EnergySource: pb.EnergySource_ENERGY_SOURCE_SOLAR, + InitTimeUtc: timestamppb.New(pivotTime.Add(time.Duration(i) * time.Hour)), + Values: yields, + } + }, + shouldErr: true, + }, + { + name: "Limit Exceeded", + setupStream: func(ctx context.Context) (pb.DataPlatformDataService_StreamCreateForecastsClient, error) { + return dc.StreamCreateForecasts(ctx) + }, + sendCount: 5001, + getReq: func(i int) *pb.CreateForecastRequest { + return &pb.CreateForecastRequest{ + LocationUuid: siteResp.LocationUuid, + Forecaster: &pb.Forecaster{ + ForecasterName: fc.ForecasterName, + ForecasterVersion: fc.ForecasterVersion, + }, + EnergySource: pb.EnergySource_ENERGY_SOURCE_SOLAR, + InitTimeUtc: timestamppb.New(pivotTime.Add(time.Duration(i) * time.Hour)), + Values: yields, + } + }, + shouldErr: true, + expectedErrCode: codes.InvalidArgument, + }, + } + + for tcIdx, tc := range testcases { + t.Run(tc.name, func(t *testing.T) { + ctx := t.Context() + + // We use a fresh site and forecaster per test case to avoid pollution + // when checking atomicity + siteRespTC := createTestLocation( + t, + fmt.Sprintf("test_stream_site_%d", tcIdx), + "POINT(-0.1 51.5)", + 1000000, + pivotTime.Add(-time.Hour*24), + nil, + ) + + fcTC := createTestForecaster(t, fmt.Sprintf("test_stream_fc_%d", tcIdx), "v1") + + stream, err := tc.setupStream(ctx) + require.NoError(t, err) + + var sendErr error + for i := 0; i < tc.sendCount; i++ { + req := tc.getReq(i) + // Overwrite the location and forecaster with the testcase-specific ones + // unless it's the deliberately broken one + if req.Forecaster.ForecasterName != "non_existent" { + req.LocationUuid = siteRespTC.LocationUuid + req.Forecaster.ForecasterName = fcTC.ForecasterName + req.Forecaster.ForecasterVersion = fcTC.ForecasterVersion + } + + if err := stream.Send(req); err != nil && err != io.EOF { + sendErr = err + break + } + } + + var ( + closeErr error + resp *pb.StreamCreateForecastsResponse + ) + + if sendErr == nil { + resp, closeErr = stream.CloseAndRecv() + } else { + closeErr = sendErr + } + + if tc.shouldErr { + require.Error(t, closeErr) + + if tc.expectedErrCode != codes.OK { + require.Equal(t, tc.expectedErrCode, status.Code(closeErr)) + } + + // Assert atomicity: No new forecasts should have been saved + postResp, err := dc.GetLatestForecasts(ctx, &pb.GetLatestForecastsRequest{ + LocationUuid: siteRespTC.LocationUuid, + EnergySource: pb.EnergySource_ENERGY_SOURCE_SOLAR, + PivotTimestampUtc: timestamppb.New( + pivotTime.Add(time.Duration(100000) * time.Hour), + ), + }) + + var postCount int + if err == nil && postResp != nil { + postCount = len(postResp.Forecasts) + } + + require.Equal(t, 0, postCount, "Atomicity failed: Forecasts were partially saved") + } else { + require.NoError(t, closeErr) + require.NotNil(t, resp) + require.Len(t, resp.ForecastUuids, tc.expectedUuidsCount) + } + }) + } +} + +func TestPrepareForecastParams(t *testing.T) { + geomID := uuid.MustParse("018e6a12-8854-7123-b123-123456789abc") + sourceID := int16(2) + forecasterID := int32(42) + + testcases := []struct { + name string + req *pb.CreateForecastRequest + expectedInitTime time.Time + expectedCreatedTime time.Time + expectDynamicCreate bool + expectedTargetLower time.Time + expectedTargetUpper time.Time + expectedResolution int16 + shouldErr bool + }{ + { + name: "Valid request with CreatedTimestampUtc", + req: &pb.CreateForecastRequest{ + InitTimeUtc: timestamppb.New( + time.Date(2024, 5, 5, 12, 30, 45, 0, time.UTC), + ), + CreatedTimestampUtc: timestamppb.New(time.Date(2024, 5, 5, 12, 0, 0, 0, time.UTC)), + Values: []*pb.CreateForecastRequest_ForecastValue{ + {HorizonMins: 30}, + {HorizonMins: 60}, + {HorizonMins: 90}, + }, + }, + expectedInitTime: time.Date( + 2024, + 5, + 5, + 12, + 30, + 0, + 0, + time.UTC, + ), // Truncated to minute + expectedCreatedTime: time.Date(2024, 5, 5, 12, 0, 0, 0, time.UTC), + expectDynamicCreate: false, + expectedTargetLower: time.Date(2024, 5, 5, 13, 0, 0, 0, time.UTC), // 12:30 + 30m + expectedTargetUpper: time.Date(2024, 5, 5, 14, 0, 0, 0, time.UTC), // 12:30 + 90m + expectedResolution: 30, + }, + { + name: "Valid request without CreatedTimestampUtc", + req: &pb.CreateForecastRequest{ + InitTimeUtc: timestamppb.New(time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC)), + Values: []*pb.CreateForecastRequest_ForecastValue{ + {HorizonMins: 0}, + {HorizonMins: 15}, + }, + }, + expectedInitTime: time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC), + expectDynamicCreate: true, // Will default to current time + expectedTargetLower: time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC), + expectedTargetUpper: time.Date(2024, 1, 1, 0, 15, 0, 0, time.UTC), + expectedResolution: 15, + }, + } + + for _, tc := range testcases { + t.Run(tc.name, func(t *testing.T) { + params, err := prepareForecastParams(tc.req, geomID, sourceID, forecasterID) + if tc.shouldErr { + require.Error(t, err) + return + } + + require.NoError(t, err) + + // Assert straightforward assignments + require.Equal(t, geomID, params.GeometryUuid) + require.Equal(t, sourceID, params.SourceTypeID) + require.Equal(t, forecasterID, params.ForecasterID) + require.Equal(t, tc.expectedResolution, params.ValueResolutionMins) + + // Assert InitTime (should be truncated to the minute) + require.Equal(t, tc.expectedInitTime, params.InitTimeUtc.Time.UTC()) + + // Assert TargetPeriod boundaries + require.True(t, params.TargetPeriod.Valid) + require.Equal(t, tc.expectedTargetLower, params.TargetPeriod.Lower.Time.UTC()) + require.Equal(t, tc.expectedTargetUpper, params.TargetPeriod.Upper.Time.UTC()) + + // Assert CreatedAt logic (either explicitly set or fallback to current time) + if tc.expectDynamicCreate { + now := time.Now().UTC().Truncate(time.Minute) + require.Equal(t, now, params.CreatedAtUtc.Time.UTC()) + } else { + require.Equal(t, tc.expectedCreatedTime, params.CreatedAtUtc.Time.UTC()) + } + + // Assert UUIDv7 timestamp encoding + uuidBytes := params.ForecastUuid + ms := uint64(uuidBytes[0])<<40 | + uint64(uuidBytes[1])<<32 | + uint64(uuidBytes[2])<<24 | + uint64(uuidBytes[3])<<16 | + uint64(uuidBytes[4])<<8 | + uint64(uuidBytes[5]) + + extractedTime := time.UnixMilli(int64(ms)).UTC() + require.Equal( + t, + tc.expectedInitTime, + extractedTime, + "UUID prefix should encode the truncated InitTimeUtc", + ) + }) + } +} diff --git a/internal/server/postgres/sql/queries/predictions.sql b/internal/server/postgres/sql/queries/predictions.sql index 18bbddc..08080af 100644 --- a/internal/server/postgres/sql/queries/predictions.sql +++ b/internal/server/postgres/sql/queries/predictions.sql @@ -63,7 +63,7 @@ ORDER BY forecaster_name ASC, created_at_utc DESC; /* --- Forecasts ------------------------------------------------------------------------------ */ --- name: CreateForecast :one +-- name: CreateForecasts :copyfrom INSERT INTO pred.forecasts ( forecast_uuid, geometry_uuid, @@ -75,30 +75,8 @@ INSERT INTO pred.forecasts ( metadata, created_at_utc ) VALUES ( - UUIDV7($4::TIMESTAMP), - $1, - $2, - $3, - $4, - $5, - TSRANGE( - $4::TIMESTAMP + MAKE_INTERVAL(mins => sqlc.arg(first_horizon_mins)::INTEGER), - $4::TIMESTAMP + MAKE_INTERVAL(mins => sqlc.arg(last_horizon_mins)::INTEGER), - '[]' - ), - CASE WHEN sqlc.arg(metadata)::JSONB = '{}'::JSONB THEN NULL ELSE sqlc.arg(metadata)::JSONB END, - CASE - WHEN sqlc.narg(created_at_utc)::TIMESTAMP IS NULL THEN CURRENT_TIMESTAMP ELSE - sqlc.narg(created_at_utc)::TIMESTAMP - END -) RETURNING - forecast_uuid, - init_time_utc, - source_type_id, - geometry_uuid, - forecaster_id, - target_period, - metadata; + $1, $2, $3, $4, $5, $6, $7, $8, $9 +); -- name: DeleteForecastByUUID :exec DELETE FROM pred.forecasts diff --git a/proto/ocf/dp/dp-data.messages.proto b/proto/ocf/dp/dp-data.messages.proto index 7be5ca0..bb9bb04 100644 --- a/proto/ocf/dp/dp-data.messages.proto +++ b/proto/ocf/dp/dp-data.messages.proto @@ -854,3 +854,8 @@ message ForecastDatum { map metadata = 9; google.protobuf.Timestamp target_timestamp_utc = 10; } + +message StreamCreateForecastsResponse { + // A list of the UUIDs generated for the successfully created forecasts. + repeated string forecast_uuids = 1; +} diff --git a/proto/ocf/dp/dp-data.service.proto b/proto/ocf/dp/dp-data.service.proto index a3bc40e..e32e653 100644 --- a/proto/ocf/dp/dp-data.service.proto +++ b/proto/ocf/dp/dp-data.service.proto @@ -82,6 +82,10 @@ service DataPlatformDataService { */ rpc StreamForecastData(StreamForecastDataRequest) returns (stream StreamForecastDataResponse) {} + /* StreamCreateForecasts allows for efficient batch creation of multiple forecasts and their values. + Note: This method is executed in a single transaction. To prevent resource exhaustion, a maximum of 5000 forecasts can be sent per stream. Exceeding this limit will abort the stream and roll back all inserts. */ + rpc StreamCreateForecasts(stream CreateForecastRequest) returns (StreamCreateForecastsResponse) {} + }