diff --git a/cmd/serve.go b/cmd/serve.go index e909518..c883f45 100644 --- a/cmd/serve.go +++ b/cmd/serve.go @@ -165,7 +165,7 @@ func runServer() error { } // Создаём веб-сервер - webServer := web.NewServer(db, logger.With("handler_type", "web")) + webServer := web.NewServer(db, logger.With("handler_type", "web"), handler.RefreshZones) httpServer := &http.Server{ Addr: webAddr, Handler: webServer, diff --git a/internal/repository/zones.go b/internal/repository/zones.go index 8b82d36..7f59e63 100644 --- a/internal/repository/zones.go +++ b/internal/repository/zones.go @@ -104,8 +104,12 @@ func (s *ZoneStorage) Update(ctx context.Context, name string, ttl, refresh, ret } affected, err := res.RowsAffected() - if affected == 0 || err != nil { - return nil, fmt.Errorf("rows not affected: %w", err) + if err != nil { + return nil, fmt.Errorf("affected rows error: %w", err) + } + + if affected == 0 { + return nil, fmt.Errorf("rows not affected") } return s.Get(ctx, name) @@ -119,7 +123,7 @@ func (s *ZoneStorage) Delete(ctx context.Context, name string) (bool, error) { affected, err := res.RowsAffected() if err != nil { - return false, err + return false, fmt.Errorf("affected rows error: %w", err) } if affected == 0 { return false, fmt.Errorf("rows not affected") diff --git a/internal/server/handler.go b/internal/server/handler.go index 7118950..53b253c 100644 --- a/internal/server/handler.go +++ b/internal/server/handler.go @@ -41,12 +41,12 @@ func NewHandler(config *HandlerConfig) *Handler { localZones: make(map[string]struct{}), logger: config.Logger, } - h.refreshZones() + h.RefreshZones() return h } -// refreshZones загружает список локальных зон из БД в in-memory кэш. -func (h *Handler) refreshZones() { +// RefreshZones загружает список локальных зон из БД в in-memory кэш. +func (h *Handler) RefreshZones() { zones, err := h.zoneStorage.GetAll(context.Background()) if err != nil { return diff --git a/internal/web/server.go b/internal/web/server.go index d1837c4..d70c15c 100644 --- a/internal/web/server.go +++ b/internal/web/server.go @@ -27,14 +27,16 @@ type Server struct { mux *http.ServeMux tmpl *template.Template logger *slog.Logger + refreshZones func() } -func NewServer(db *sql.DB, logger *slog.Logger) *Server { +func NewServer(db *sql.DB, logger *slog.Logger, refeshZones func()) *Server { s := &Server{ zoneStorage: repository.NewZoneStorage(db), recordStorage: repository.NewRecordStorage(db), mux: http.NewServeMux(), logger: logger, + refreshZones: refeshZones, } // Парсим шаблоны из embed.FS @@ -53,7 +55,11 @@ func NewServer(db *sql.DB, logger *slog.Logger) *Server { s.mux.HandleFunc("DELETE /api/records/{id}", s.handleDeleteRecord) // Static files - staticSub, _ := fs.Sub(staticFS, "static") + staticSub, err := fs.Sub(staticFS, "static") + if err != nil { + panic(err) + } + s.mux.Handle("GET /static/", http.StripPrefix("/static/", http.FileServer(http.FS(staticSub)))) // SPA — все остальные пути отдаём index.html @@ -134,6 +140,7 @@ func (s *Server) handleCreateZone(w http.ResponseWriter, r *http.Request) { writeJSON(w, http.StatusConflict, map[string]string{"error": err.Error()}) return } + s.refreshZones() writeJSON(w, http.StatusCreated, zone) } @@ -176,6 +183,7 @@ func (s *Server) handleDeleteZone(w http.ResponseWriter, r *http.Request) { writeJSON(w, http.StatusInternalServerError, map[string]string{"error": err.Error()}) return } + s.refreshZones() w.WriteHeader(http.StatusNoContent) } diff --git a/internal/web/server_test.go b/internal/web/server_test.go index 8c78b75..07b8b64 100644 --- a/internal/web/server_test.go +++ b/internal/web/server_test.go @@ -47,10 +47,11 @@ func upDatabase(t *testing.T) *sql.DB { return conn } -func setupTest(t *testing.T) (*Server, *sql.DB) { +func setupTest(t *testing.T) (*Server, *sql.DB, func() int) { + updateZones := 0 db := upDatabase(t) - s := NewServer(db, slog.New(slog.NewTextHandler(io.Discard, nil))) - return s, db + s := NewServer(db, slog.New(slog.NewTextHandler(io.Discard, nil)), func() { updateZones++ }) + return s, db, func() int { return updateZones } } func seedZone(t *testing.T, db *sql.DB, name string, ttl, refresh, retry, expire int64) *repository.Zone { @@ -67,7 +68,7 @@ func seedRecord(t *testing.T, db *sql.DB, zoneID uint64, domain, rtype, rdata st return rec } -func request(t *testing.T, s *Server, method, path, body string) *httptest.ResponseRecorder { +func request(_ *testing.T, s *Server, method, path, body string) *httptest.ResponseRecorder { var reqBody *strings.Reader if body != "" { reqBody = strings.NewReader(body) @@ -85,18 +86,19 @@ func request(t *testing.T, s *Server, method, path, body string) *httptest.Respo // --- Zones --- func TestWeb_ListZones_Empty(t *testing.T) { - s, db := setupTest(t) + s, db, getUpdatedZones := setupTest(t) defer db.Close() w := request(t, s, "GET", "/api/zones", "") assert.Equal(t, http.StatusOK, w.Code) + assert.Zero(t, getUpdatedZones()) assert.Equal(t, "application/json; charset=utf-8", w.Header().Get("Content-Type")) assert.Equal(t, "null\n", w.Body.String()) } func TestWeb_ListZones(t *testing.T) { - s, db := setupTest(t) + s, db, getUpdatedZones := setupTest(t) defer db.Close() seedZone(t, db, "example.com.", 300, 3600, 300, 86400) @@ -124,10 +126,11 @@ func TestWeb_ListZones(t *testing.T) { assert.Equal(t, int64(86401), zones[1].Expire) assert.NotZero(t, zones[1].Serial) + assert.Zero(t, getUpdatedZones()) } func TestWeb_GetZone(t *testing.T) { - s, db := setupTest(t) + s, db, getZoneUpdates := setupTest(t) defer db.Close() seedZone(t, db, "example.com.", 300, 3600, 400, 86400) @@ -144,11 +147,11 @@ func TestWeb_GetZone(t *testing.T) { assert.Equal(t, int64(3600), zone.Refresh) assert.Equal(t, int64(400), zone.Retry) assert.Equal(t, int64(86400), zone.Expire) - + assert.Zero(t, getZoneUpdates()) } func TestWeb_GetZone_NotFound(t *testing.T) { - s, db := setupTest(t) + s, db, _ := setupTest(t) defer db.Close() w := request(t, s, "GET", "/api/zones/nonexistent.com.", "") @@ -162,7 +165,7 @@ func TestWeb_GetZone_NotFound(t *testing.T) { } func TestWeb_CreateZone(t *testing.T) { - s, db := setupTest(t) + s, db, getZoneUpdates := setupTest(t) defer db.Close() w := request(t, s, "POST", "/api/zones", `{"name":"example.com.","ttl":3600}`) @@ -175,15 +178,17 @@ func TestWeb_CreateZone(t *testing.T) { assert.NotZero(t, zone.Id) assert.Equal(t, "example.com.", zone.Name) assert.Equal(t, int64(3600), zone.TTL) + assert.Equal(t, 1, getZoneUpdates()) } func TestWeb_CreateZone_AddsDot(t *testing.T) { - s, db := setupTest(t) + s, db, getZoneUpdates := setupTest(t) defer db.Close() w := request(t, s, "POST", "/api/zones", `{"name":"example.com","ttl":300}`) assert.Equal(t, http.StatusCreated, w.Code) + assert.Equal(t, 1, getZoneUpdates()) var zone repository.Zone err := json.Unmarshal(w.Body.Bytes(), &zone) @@ -192,25 +197,27 @@ func TestWeb_CreateZone_AddsDot(t *testing.T) { } func TestWeb_CreateZone_EmptyName(t *testing.T) { - s, db := setupTest(t) + s, db, getZoneUpdates := setupTest(t) defer db.Close() w := request(t, s, "POST", "/api/zones", `{"name":"","ttl":300}`) + assert.Zero(t, getZoneUpdates()) assert.Equal(t, http.StatusBadRequest, w.Code) } func TestWeb_CreateZone_InvalidJSON(t *testing.T) { - s, db := setupTest(t) + s, db, getZoneUpdates := setupTest(t) defer db.Close() w := request(t, s, "POST", "/api/zones", `not json`) assert.Equal(t, http.StatusBadRequest, w.Code) + assert.Zero(t, getZoneUpdates()) } func TestWeb_CreateZone_Duplicate(t *testing.T) { - s, db := setupTest(t) + s, db, getZoneUpdates := setupTest(t) defer db.Close() seedZone(t, db, "example.com.", 300, 3600, 400, 8640000) @@ -218,10 +225,11 @@ func TestWeb_CreateZone_Duplicate(t *testing.T) { w := request(t, s, "POST", "/api/zones", `{"name":"example.com.","ttl":300,"refresh":3600,"retry":400,"expire":8640000}`) assert.Equal(t, http.StatusConflict, w.Code) + assert.Zero(t, getZoneUpdates()) } func TestWeb_UpdateZone(t *testing.T) { - s, db := setupTest(t) + s, db, getZoneUpdates := setupTest(t) defer db.Close() seedZone(t, db, "example.com.", 300, 3600, 400, 86400) @@ -229,6 +237,7 @@ func TestWeb_UpdateZone(t *testing.T) { w := request(t, s, "PUT", "/api/zones/example.com.", `{"ttl":600,"refresh":3601,"retry":401,"expire":8640001}`) assert.Equal(t, http.StatusOK, w.Code) + assert.Zero(t, getZoneUpdates()) var zone repository.Zone err := json.Unmarshal(w.Body.Bytes(), &zone) @@ -242,7 +251,7 @@ func TestWeb_UpdateZone(t *testing.T) { } func TestWeb_UpdateZone_NotFound(t *testing.T) { - s, db := setupTest(t) + s, db, _ := setupTest(t) defer db.Close() w := request(t, s, "PUT", "/api/zones/nonexistent.com.", `{"ttl":600,"refresh":3601,"retry":401,"expire":8640001}`) @@ -251,7 +260,7 @@ func TestWeb_UpdateZone_NotFound(t *testing.T) { } func TestWeb_DeleteZone(t *testing.T) { - s, db := setupTest(t) + s, db, getZoneUpdates := setupTest(t) defer db.Close() seedZone(t, db, "example.com.", 300, 3600, 400, 86400) @@ -259,6 +268,7 @@ func TestWeb_DeleteZone(t *testing.T) { w := request(t, s, "DELETE", "/api/zones/example.com.", "") assert.Equal(t, http.StatusNoContent, w.Code) + assert.Equal(t, 1, getZoneUpdates()) // Проверяем, что зона удалена w2 := request(t, s, "GET", "/api/zones/example.com.", "") @@ -266,18 +276,19 @@ func TestWeb_DeleteZone(t *testing.T) { } func TestWeb_DeleteZone_NotFound(t *testing.T) { - s, db := setupTest(t) + s, db, getZoneUpdates := setupTest(t) defer db.Close() w := request(t, s, "DELETE", "/api/zones/nonexistent.com.", "") assert.Equal(t, http.StatusInternalServerError, w.Code) + assert.Zero(t, getZoneUpdates()) } // --- Records --- func TestWeb_ListRecords_Empty(t *testing.T) { - s, db := setupTest(t) + s, db, _ := setupTest(t) defer db.Close() seedZone(t, db, "example.com.", 300, 3600, 400, 86400) @@ -289,7 +300,7 @@ func TestWeb_ListRecords_Empty(t *testing.T) { } func TestWeb_ListRecords(t *testing.T) { - s, db := setupTest(t) + s, db, _ := setupTest(t) defer db.Close() zone := seedZone(t, db, "example.com.", 300, 3600, 400, 86400) @@ -311,7 +322,7 @@ func TestWeb_ListRecords(t *testing.T) { } func TestWeb_ListRecords_ZoneNotFound(t *testing.T) { - s, db := setupTest(t) + s, db, _ := setupTest(t) defer db.Close() w := request(t, s, "GET", "/api/zones/nonexistent.com./records", "") @@ -320,7 +331,7 @@ func TestWeb_ListRecords_ZoneNotFound(t *testing.T) { } func TestWeb_CreateRecord(t *testing.T) { - s, db := setupTest(t) + s, db, _ := setupTest(t) defer db.Close() seedZone(t, db, "example.com.", 300, 3600, 400, 86400) @@ -341,7 +352,7 @@ func TestWeb_CreateRecord(t *testing.T) { } func TestWeb_CreateRecord_DefaultTTL(t *testing.T) { - s, db := setupTest(t) + s, db, _ := setupTest(t) defer db.Close() seedZone(t, db, "example.com.", 600, 3600, 400, 86400) @@ -359,7 +370,7 @@ func TestWeb_CreateRecord_DefaultTTL(t *testing.T) { } func TestWeb_CreateRecord_ZoneNotFound(t *testing.T) { - s, db := setupTest(t) + s, db, _ := setupTest(t) defer db.Close() w := request(t, s, "POST", "/api/zones/nonexistent.com./records", @@ -369,7 +380,7 @@ func TestWeb_CreateRecord_ZoneNotFound(t *testing.T) { } func TestWeb_CreateRecord_MissingFields(t *testing.T) { - s, db := setupTest(t) + s, db, _ := setupTest(t) defer db.Close() seedZone(t, db, "example.com.", 300, 3600, 400, 86400) @@ -381,7 +392,7 @@ func TestWeb_CreateRecord_MissingFields(t *testing.T) { } func TestWeb_UpdateRecord(t *testing.T) { - s, db := setupTest(t) + s, db, _ := setupTest(t) defer db.Close() zone := seedZone(t, db, "example.com.", 300, 3600, 400, 86400) @@ -405,7 +416,7 @@ func TestWeb_UpdateRecord(t *testing.T) { } func TestWeb_UpdateRecord_NotFound(t *testing.T) { - s, db := setupTest(t) + s, db, _ := setupTest(t) defer db.Close() zone := seedZone(t, db, "example.com.", 300, 3600, 400, 86400) @@ -420,7 +431,7 @@ func TestWeb_UpdateRecord_NotFound(t *testing.T) { } func TestWeb_UpdateRecord_InvalidID(t *testing.T) { - s, db := setupTest(t) + s, db, _ := setupTest(t) defer db.Close() w := request(t, s, "PUT", "/api/records/abc", @@ -430,7 +441,7 @@ func TestWeb_UpdateRecord_InvalidID(t *testing.T) { } func TestWeb_DeleteRecord(t *testing.T) { - s, db := setupTest(t) + s, db, _ := setupTest(t) defer db.Close() zone := seedZone(t, db, "example.com.", 300, 3600, 400, 86400) @@ -442,7 +453,7 @@ func TestWeb_DeleteRecord(t *testing.T) { } func TestWeb_DeleteRecord_NotFound(t *testing.T) { - s, db := setupTest(t) + s, db, _ := setupTest(t) defer db.Close() w := request(t, s, "DELETE", "/api/records/999", "") @@ -451,7 +462,7 @@ func TestWeb_DeleteRecord_NotFound(t *testing.T) { } func TestWeb_DeleteRecord_InvalidID(t *testing.T) { - s, db := setupTest(t) + s, db, _ := setupTest(t) defer db.Close() w := request(t, s, "DELETE", "/api/records/abc", "") @@ -462,7 +473,7 @@ func TestWeb_DeleteRecord_InvalidID(t *testing.T) { // --- Static files --- func TestWeb_IndexPage(t *testing.T) { - s, db := setupTest(t) + s, db, _ := setupTest(t) defer db.Close() w := request(t, s, "GET", "/", "") @@ -473,7 +484,7 @@ func TestWeb_IndexPage(t *testing.T) { } func TestWeb_StaticCSS(t *testing.T) { - s, db := setupTest(t) + s, db, _ := setupTest(t) defer db.Close() w := request(t, s, "GET", "/static/style.css", "") @@ -484,7 +495,7 @@ func TestWeb_StaticCSS(t *testing.T) { } func TestWeb_StaticJS(t *testing.T) { - s, db := setupTest(t) + s, db, _ := setupTest(t) defer db.Close() w := request(t, s, "GET", "/static/app.js", "") @@ -495,7 +506,7 @@ func TestWeb_StaticJS(t *testing.T) { } func TestWeb_NotFound(t *testing.T) { - s, db := setupTest(t) + s, db, _ := setupTest(t) defer db.Close() w := request(t, s, "GET", "/nonexistent", "")