Skip to content

Commit e9db206

Browse files
committed
Add context.Context to ProviderUsecase and DomainUsecase interfaces
Propagate context.Context as first parameter through all provider and domain usecase interface methods that didn't already have it. This is a prerequisite for the upcoming secret management layer, which needs request-scoped context to carry session-derived encryption keys.
1 parent af51790 commit e9db206

19 files changed

Lines changed: 130 additions & 120 deletions

internal/api-admin/controller/provider_controller.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -184,7 +184,7 @@ func (pc *ProviderController) UpdateProvider(c *gin.Context) {
184184
func (pc *ProviderController) ClearProviders(c *gin.Context) {
185185
user := middleware.MyUser(c)
186186
if user != nil {
187-
providers, err := pc.providerService.ListUserProviders(user)
187+
providers, err := pc.providerService.ListUserProviders(c.Request.Context(), user)
188188
if err != nil {
189189
middleware.ErrorResponse(c, http.StatusInternalServerError, err)
190190
return

internal/api/controller/domain.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -106,7 +106,7 @@ func (dc *DomainController) AddDomain(c *gin.Context) {
106106
return
107107
}
108108

109-
err = dc.domainService.CreateDomain(user, &uz)
109+
err = dc.domainService.CreateDomain(c.Request.Context(), user, &uz)
110110
if err != nil {
111111
middleware.ErrorResponse(c, http.StatusInternalServerError, err)
112112
return

internal/api/controller/provider.go

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -61,7 +61,7 @@ func (pc *ProviderController) ListProviders(c *gin.Context) {
6161
return
6262
}
6363

64-
providers, err := pc.providerService.ListUserProviders(user)
64+
providers, err := pc.providerService.ListUserProviders(c.Request.Context(), user)
6565
if err != nil {
6666
middleware.ErrorResponse(c, http.StatusInternalServerError, err)
6767
return
@@ -119,7 +119,7 @@ func (pc *ProviderController) AddProvider(c *gin.Context) {
119119
return
120120
}
121121

122-
provider, err := pc.providerService.CreateProvider(user, &usrc)
122+
provider, err := pc.providerService.CreateProvider(c.Request.Context(), user, &usrc)
123123
if err != nil {
124124
middleware.ErrorResponse(c, http.StatusInternalServerError, err)
125125
return
@@ -157,7 +157,7 @@ func (pc *ProviderController) UpdateProvider(c *gin.Context) {
157157
return
158158
}
159159

160-
err = pc.providerService.UpdateProviderFromMessage(old.Id, user, &provider)
160+
err = pc.providerService.UpdateProviderFromMessage(c.Request.Context(), old.Id, user, &provider)
161161
if err != nil {
162162
middleware.ErrorResponse(c, http.StatusInternalServerError, err)
163163
return
@@ -191,7 +191,7 @@ func (pc *ProviderController) DeleteProvider(c *gin.Context) {
191191

192192
providermeta := c.MustGet("providermeta").(*happydns.ProviderMeta)
193193

194-
err := pc.providerService.DeleteProvider(user, providermeta.Id)
194+
err := pc.providerService.DeleteProvider(c.Request.Context(), user, providermeta.Id)
195195
if err != nil {
196196
middleware.ErrorResponse(c, http.StatusInternalServerError, err)
197197
return
@@ -220,7 +220,7 @@ func (pc *ProviderController) DeleteProvider(c *gin.Context) {
220220
func (pc *ProviderController) GetDomainsHostedByProvider(c *gin.Context) {
221221
provider := c.MustGet("provider").(*happydns.Provider)
222222

223-
domains, err := pc.providerService.ListHostedDomains(provider)
223+
domains, err := pc.providerService.ListHostedDomains(c.Request.Context(), provider)
224224
if err != nil {
225225
middleware.ErrorResponse(c, http.StatusBadRequest, err)
226226
return
@@ -249,7 +249,7 @@ func (pc *ProviderController) GetDomainsHostedByProvider(c *gin.Context) {
249249
func (pc *ProviderController) CreateDomainOnProvider(c *gin.Context) {
250250
provider := c.MustGet("provider").(*happydns.Provider)
251251

252-
err := pc.providerService.CreateDomainOnProvider(provider, c.Param("fqdn"))
252+
err := pc.providerService.CreateDomainOnProvider(c.Request.Context(), provider, c.Param("fqdn"))
253253
if err != nil {
254254
middleware.ErrorResponse(c, http.StatusBadRequest, err)
255255
return

internal/api/controller/provider_settings.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -78,7 +78,7 @@ func (psc *ProviderSettingsController) NextProviderSettingsState(c *gin.Context)
7878
return
7979
}
8080

81-
provider, form, err := psc.pSettingsServices.NextProviderSettingsState(&uss, pType, user)
81+
provider, form, err := psc.pSettingsServices.NextProviderSettingsState(c.Request.Context(), &uss, pType, user)
8282
if err != nil {
8383
middleware.ErrorResponse(c, http.StatusInternalServerError, err)
8484
return

internal/api/middleware/provider.go

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -47,7 +47,7 @@ func ProviderMetaHandler(providerService happydns.ProviderUsecase) gin.HandlerFu
4747
}
4848

4949
// Retrieve provider meta
50-
providermeta, err := providerService.GetUserProviderMeta(user, pid)
50+
providermeta, err := providerService.GetUserProviderMeta(c.Request.Context(), user, pid)
5151
if err != nil {
5252
ErrorResponse(c, http.StatusNotFound, fmt.Errorf("provider not found"))
5353
return
@@ -77,7 +77,7 @@ func ProviderHandler(providerService happydns.ProviderUsecase) gin.HandlerFunc {
7777
}
7878

7979
// Retrieve provider
80-
provider, err := providerService.GetUserProvider(user, pid)
80+
provider, err := providerService.GetUserProvider(c.Request.Context(), user, pid)
8181
if err != nil {
8282
ErrorResponse(c, http.StatusNotFound, fmt.Errorf("provider not found"))
8383
return

internal/usecase/domain/domain.go

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,7 @@
2222
package domain
2323

2424
import (
25+
"context"
2526
"errors"
2627
"fmt"
2728

@@ -32,12 +33,12 @@ import (
3233

3334
// ProviderGetter is an interface for getting providers.
3435
type ProviderGetter interface {
35-
GetUserProvider(user *happydns.User, providerID happydns.Identifier) (*happydns.Provider, error)
36+
GetUserProvider(ctx context.Context, user *happydns.User, providerID happydns.Identifier) (*happydns.Provider, error)
3637
}
3738

3839
// DomainExistenceTester is an interface for testing domain existence.
3940
type DomainExistenceTester interface {
40-
TestDomainExistence(provider *happydns.Provider, name string) error
41+
TestDomainExistence(ctx context.Context, provider *happydns.Provider, name string) error
4142
}
4243

4344
type Service struct {
@@ -65,18 +66,18 @@ func NewService(
6566
}
6667

6768
// CreateDomain creates a new domain for the given user.
68-
func (s *Service) CreateDomain(user *happydns.User, uz *happydns.Domain) error {
69+
func (s *Service) CreateDomain(ctx context.Context, user *happydns.User, uz *happydns.Domain) error {
6970
uz, err := happydns.NewDomain(user, uz.DomainName, uz.ProviderId)
7071
if err != nil {
7172
return err
7273
}
7374

74-
provider, err := s.providerService.GetUserProvider(user, uz.ProviderId)
75+
provider, err := s.providerService.GetUserProvider(ctx, user, uz.ProviderId)
7576
if err != nil {
7677
return happydns.ValidationError{Msg: fmt.Sprintf("unable to find the provider.")}
7778
}
7879

79-
if err = s.domainExistence.TestDomainExistence(provider, uz.DomainName); err != nil {
80+
if err = s.domainExistence.TestDomainExistence(ctx, provider, uz.DomainName); err != nil {
8081
return happydns.NotFoundError{Msg: err.Error()}
8182
}
8283

internal/usecase/domain/domain_test.go

Lines changed: 16 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,7 @@
2222
package domain_test
2323

2424
import (
25+
"context"
2526
"fmt"
2627
"testing"
2728

@@ -35,6 +36,8 @@ import (
3536
"git.happydns.org/happyDomain/model"
3637
)
3738

39+
var ctx = context.Background()
40+
3841
// Mock implementations for testing
3942

4043
func init() {
@@ -159,7 +162,7 @@ func Test_CreateDomain(t *testing.T) {
159162
ProviderId: providerId,
160163
}
161164

162-
err := service.CreateDomain(user, domainToCreate)
165+
err := service.CreateDomain(ctx, user, domainToCreate)
163166
if err != nil {
164167
t.Fatalf("unexpected error: %v", err)
165168
}
@@ -205,7 +208,7 @@ func Test_CreateDomain_InvalidProvider(t *testing.T) {
205208
ProviderId: invalidProviderId,
206209
}
207210

208-
err := service.CreateDomain(user, domainToCreate)
211+
err := service.CreateDomain(ctx, user, domainToCreate)
209212
if err == nil {
210213
t.Error("expected error when creating domain with invalid provider")
211214
}
@@ -224,7 +227,7 @@ func Test_GetUserDomain(t *testing.T) {
224227
DomainName: "example.com",
225228
ProviderId: providerId,
226229
}
227-
err := service.CreateDomain(user, domainToCreate)
230+
err := service.CreateDomain(ctx, user, domainToCreate)
228231
if err != nil {
229232
t.Fatalf("failed to create domain: %v", err)
230233
}
@@ -265,7 +268,7 @@ func Test_GetUserDomain_WrongUser(t *testing.T) {
265268
DomainName: "user1-domain.com",
266269
ProviderId: providerId,
267270
}
268-
err := service.CreateDomain(user1, domainToCreate)
271+
err := service.CreateDomain(ctx, user1, domainToCreate)
269272
if err != nil {
270273
t.Fatalf("failed to create domain: %v", err)
271274
}
@@ -317,7 +320,7 @@ func Test_GetUserDomainByFQDN(t *testing.T) {
317320
DomainName: "example.com.",
318321
ProviderId: providerId,
319322
}
320-
err := service.CreateDomain(user, domainToCreate)
323+
err := service.CreateDomain(ctx, user, domainToCreate)
321324
if err != nil {
322325
t.Fatalf("failed to create domain: %v", err)
323326
}
@@ -352,7 +355,7 @@ func Test_ListUserDomains(t *testing.T) {
352355
DomainName: name,
353356
ProviderId: providerId,
354357
}
355-
err := service.CreateDomain(user, domainToCreate)
358+
err := service.CreateDomain(ctx, user, domainToCreate)
356359
if err != nil {
357360
t.Fatalf("failed to create domain %s: %v", name, err)
358361
}
@@ -385,7 +388,7 @@ func Test_ListUserDomains_MultipleUsers(t *testing.T) {
385388
DomainName: fmt.Sprintf("user1-domain%d.com", i),
386389
ProviderId: providerId1,
387390
}
388-
err := service.CreateDomain(user1, domainToCreate)
391+
err := service.CreateDomain(ctx, user1, domainToCreate)
389392
if err != nil {
390393
t.Fatalf("failed to create domain: %v", err)
391394
}
@@ -396,7 +399,7 @@ func Test_ListUserDomains_MultipleUsers(t *testing.T) {
396399
DomainName: "user2-domain.com",
397400
ProviderId: providerId2,
398401
}
399-
err := service.CreateDomain(user2, domainToCreate)
402+
err := service.CreateDomain(ctx, user2, domainToCreate)
400403
if err != nil {
401404
t.Fatalf("failed to create domain: %v", err)
402405
}
@@ -451,7 +454,7 @@ func Test_UpdateDomain(t *testing.T) {
451454
DomainName: "example.com",
452455
ProviderId: providerId,
453456
}
454-
err := service.CreateDomain(user, domainToCreate)
457+
err := service.CreateDomain(ctx, user, domainToCreate)
455458
if err != nil {
456459
t.Fatalf("failed to create domain: %v", err)
457460
}
@@ -503,7 +506,7 @@ func Test_UpdateDomain_PreventIdChange(t *testing.T) {
503506
DomainName: "example.com",
504507
ProviderId: providerId,
505508
}
506-
err := service.CreateDomain(user, domainToCreate)
509+
err := service.CreateDomain(ctx, user, domainToCreate)
507510
if err != nil {
508511
t.Fatalf("failed to create domain: %v", err)
509512
}
@@ -546,7 +549,7 @@ func Test_UpdateDomain_WrongUser(t *testing.T) {
546549
DomainName: "user1-domain.com",
547550
ProviderId: providerId,
548551
}
549-
err := service.CreateDomain(user1, domainToCreate)
552+
err := service.CreateDomain(ctx, user1, domainToCreate)
550553
if err != nil {
551554
t.Fatalf("failed to create domain: %v", err)
552555
}
@@ -580,7 +583,7 @@ func Test_DeleteDomain(t *testing.T) {
580583
DomainName: "example.com",
581584
ProviderId: providerId,
582585
}
583-
err := service.CreateDomain(user, domainToCreate)
586+
err := service.CreateDomain(ctx, user, domainToCreate)
584587
if err != nil {
585588
t.Fatalf("failed to create domain: %v", err)
586589
}
@@ -621,7 +624,7 @@ func Test_UpdateDomain_Alias(t *testing.T) {
621624
DomainName: "example.com",
622625
ProviderId: providerId,
623626
}
624-
err := service.CreateDomain(user, domainToCreate)
627+
err := service.CreateDomain(ctx, user, domainToCreate)
625628
if err != nil {
626629
t.Fatalf("failed to create domain: %v", err)
627630
}

internal/usecase/orchestrator/factory.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -41,7 +41,7 @@ type DomainUpdater interface {
4141

4242
// ProviderGetter is an interface for getting providers.
4343
type ProviderGetter interface {
44-
GetUserProvider(user *happydns.User, providerID happydns.Identifier) (*happydns.Provider, error)
44+
GetUserProvider(ctx context.Context, user *happydns.User, providerID happydns.Identifier) (*happydns.Provider, error)
4545
}
4646

4747
// ZoneRetriever is an interface for retrieving zones from providers.

internal/usecase/orchestrator/remote_zone_importer.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -60,7 +60,7 @@ func NewRemoteZoneImporterUsecase(
6060
// and imports them via ZoneImporterUsecase. A domain log entry is appended on
6161
// success. Returns the newly created zone or an error.
6262
func (uc *RemoteZoneImporterUsecase) Import(ctx context.Context, user *happydns.User, domain *happydns.Domain) (*happydns.Zone, error) {
63-
provider, err := uc.providerService.GetUserProvider(user, domain.ProviderId)
63+
provider, err := uc.providerService.GetUserProvider(ctx, user, domain.ProviderId)
6464
if err != nil {
6565
return nil, err
6666
}

internal/usecase/orchestrator/zone_correction_applier.go

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -94,7 +94,7 @@ func (uc *ZoneCorrectionApplierUsecase) computeExecutableCorrections(
9494
targetRecords = adapter.BuildTargetRecords(providerRecords, corrections, wantedCorrections)
9595

9696
// Step 3: Get executable corrections from the provider for the target state.
97-
provider, err := uc.providerService.GetUserProvider(user, domain.ProviderId)
97+
provider, err := uc.providerService.GetUserProvider(ctx, user, domain.ProviderId)
9898
if err != nil {
9999
return nil, nil, nil, nbDiffs, err
100100
}
@@ -179,7 +179,7 @@ func (uc *ZoneCorrectionApplierUsecase) Apply(
179179
// Step 4b: If provider manages SOA serial, re-fetch to get the actual published state.
180180
publishedRecords := targetRecords
181181
refetched := false
182-
provider, provErr := uc.providerService.GetUserProvider(user, domain.ProviderId)
182+
provider, provErr := uc.providerService.GetUserProvider(ctx, user, domain.ProviderId)
183183
if provErr == nil && providerReg.ProviderHasCapability(provider, "manages-soa-serial") {
184184
fetched, fetchErr := uc.zoneRetriever.RetrieveZone(ctx, provider, domain.DomainName)
185185
if fetchErr != nil {

0 commit comments

Comments
 (0)