package server import ( "encoding/json" "errors" "net/http" "strconv" "github.com/go-chi/chi/v5" "git.unkin.net/unkin/tomswallapi/internal/compiler" "git.unkin.net/unkin/tomswallapi/internal/model" "git.unkin.net/unkin/tomswallapi/internal/store" ) // mountResources wires the Terraform-facing CRUD endpoints. Resources with a // dedicated repository method are wired here; the long-tail per-device sections // (providers, tc, etc.) are added as their storage lands. func (s *Server) mountResources(r chi.Router) { r.Get("/generation", s.handleGeneration) r.Route("/fabrics", func(r chi.Router) { r.Get("/", s.listFabrics) r.Get("/{name}", s.getFabric) r.Put("/{name}", s.putFabric) r.Delete("/{name}", s.deleteFabric) }) r.Route("/zones", func(r chi.Router) { r.Get("/", s.listZones) r.Get("/{name}", s.getZone) r.Put("/{name}", s.putZone) r.Delete("/{name}", s.deleteZone) }) r.Route("/address-groups", func(r chi.Router) { r.Get("/", s.listAddressGroups) r.Get("/{name}", s.getAddressGroup) r.Put("/{name}", s.putAddressGroup) r.Delete("/{name}", s.deleteAddressGroup) }) r.Route("/devices", func(r chi.Router) { r.Get("/", s.listDevices) r.Get("/{name}", s.getDevice) r.Put("/{name}", s.putDevice) r.Delete("/{name}", s.deleteDevice) r.Get("/{name}/bindings", s.listBindings) r.Get("/{name}/bindings/{zone}", s.getBinding) r.Put("/{name}/bindings/{zone}", s.putBinding) r.Delete("/{name}/bindings/{zone}", s.deleteBinding) }) r.Route("/portgroups", func(r chi.Router) { r.Get("/", s.listPortGroups) r.Get("/{name}", s.getPortGroup) r.Put("/{name}", s.putPortGroup) r.Delete("/{name}", s.deletePortGroup) }) r.Route("/rules", func(r chi.Router) { r.Get("/", s.listRules) r.Post("/", s.createRule) r.Get("/{id}", s.getRule) r.Delete("/{id}", s.deleteRule) }) s.mountNAT(r) } // respondOne writes a single resource, mapping ErrNotFound to 404. func respondOne(w http.ResponseWriter, v any, err error) { if err != nil { if errors.Is(err, store.ErrNotFound) { writeError(w, http.StatusNotFound, "not found") return } writeError(w, http.StatusInternalServerError, err.Error()) return } writeJSON(w, http.StatusOK, v) } // respondDelete maps a delete result to 204/404/500. func respondDelete(w http.ResponseWriter, err error) { if err != nil { if errors.Is(err, store.ErrNotFound) { writeError(w, http.StatusNotFound, "not found") return } writeError(w, http.StatusInternalServerError, err.Error()) return } w.WriteHeader(http.StatusNoContent) } func (s *Server) listPortGroups(w http.ResponseWriter, r *http.Request) { list, err := s.store.ListPortGroups(r.Context()) respondList(w, list, err) } func (s *Server) putPortGroup(w http.ResponseWriter, r *http.Request) { var p model.PortGroup if !decode(w, r, &p) { return } p.Name = chi.URLParam(r, "name") if err := s.store.UpsertPortGroup(r.Context(), p); err != nil { writeError(w, http.StatusInternalServerError, err.Error()) return } writeJSON(w, http.StatusOK, p) } func (s *Server) handleGeneration(w http.ResponseWriter, r *http.Request) { g, err := s.store.Generation(r.Context()) if err != nil { writeError(w, http.StatusInternalServerError, err.Error()) return } writeJSON(w, http.StatusOK, map[string]int64{"generation": g}) } // ---- Fabrics --------------------------------------------------------------- func (s *Server) listFabrics(w http.ResponseWriter, r *http.Request) { list, err := s.store.ListFabrics(r.Context()) respondList(w, list, err) } func (s *Server) putFabric(w http.ResponseWriter, r *http.Request) { var f model.Fabric if !decode(w, r, &f) { return } f.Name = chi.URLParam(r, "name") if err := s.store.UpsertFabric(r.Context(), f); err != nil { writeError(w, http.StatusInternalServerError, err.Error()) return } writeJSON(w, http.StatusOK, f) } // ---- Zones ----------------------------------------------------------------- func (s *Server) listZones(w http.ResponseWriter, r *http.Request) { list, err := s.store.ListZones(r.Context()) respondList(w, list, err) } func (s *Server) putZone(w http.ResponseWriter, r *http.Request) { var z model.Zone if !decode(w, r, &z) { return } z.Name = chi.URLParam(r, "name") if err := s.store.UpsertZone(r.Context(), z); err != nil { writeError(w, http.StatusInternalServerError, err.Error()) return } writeJSON(w, http.StatusOK, z) } // ---- Address groups -------------------------------------------------------- func (s *Server) listAddressGroups(w http.ResponseWriter, r *http.Request) { list, err := s.store.ListAddressGroups(r.Context()) respondList(w, list, err) } func (s *Server) putAddressGroup(w http.ResponseWriter, r *http.Request) { var g model.AddressGroup if !decode(w, r, &g) { return } g.Name = chi.URLParam(r, "name") if g.Type != model.GroupStatic && g.Type != model.GroupDNS && g.Type != model.GroupASN { writeError(w, http.StatusBadRequest, "type must be one of: static, dns, asn") return } if err := s.store.UpsertAddressGroup(r.Context(), g); err != nil { writeError(w, http.StatusInternalServerError, err.Error()) return } writeJSON(w, http.StatusOK, g) } // ---- Devices & bindings ---------------------------------------------------- func (s *Server) listDevices(w http.ResponseWriter, r *http.Request) { list, err := s.store.ListDevices(r.Context()) respondList(w, list, err) } func (s *Server) putDevice(w http.ResponseWriter, r *http.Request) { var d model.Device if !decode(w, r, &d) { return } d.Name = chi.URLParam(r, "name") if d.Class != model.ClassRouter && d.Class != model.ClassFirewall { writeError(w, http.StatusBadRequest, "class must be one of: router, firewall") return } if err := s.store.UpsertDevice(r.Context(), d); err != nil { writeError(w, http.StatusInternalServerError, err.Error()) return } writeJSON(w, http.StatusOK, d) } func (s *Server) listBindings(w http.ResponseWriter, r *http.Request) { list, err := s.store.ListBindings(r.Context(), chi.URLParam(r, "name")) respondList(w, list, err) } func (s *Server) putBinding(w http.ResponseWriter, r *http.Request) { var b model.Binding if !decode(w, r, &b) { return } b.Device = chi.URLParam(r, "name") b.Zone = chi.URLParam(r, "zone") if err := s.store.UpsertBinding(r.Context(), b); err != nil { writeError(w, http.StatusInternalServerError, err.Error()) return } writeJSON(w, http.StatusOK, b) } // ---- Rules ----------------------------------------------------------------- func (s *Server) listRules(w http.ResponseWriter, r *http.Request) { list, err := s.store.ListRules(r.Context()) respondList(w, list, err) } func (s *Server) createRule(w http.ResponseWriter, r *http.Request) { var rule model.Rule if !decode(w, r, &rule) { return } id, err := s.store.CreateRule(r.Context(), rule) if err != nil { // Grammar/validation failures are client errors. writeError(w, http.StatusBadRequest, err.Error()) return } rule.ID = id writeJSON(w, http.StatusCreated, rule) } func (s *Server) deleteRule(w http.ResponseWriter, r *http.Request) { id, err := strconv.ParseInt(chi.URLParam(r, "id"), 10, 64) if err != nil { writeError(w, http.StatusBadRequest, "id must be an integer") return } if err := s.store.DeleteRule(r.Context(), id); err != nil { if errors.Is(err, store.ErrNotFound) { writeError(w, http.StatusNotFound, "rule not found") return } writeError(w, http.StatusInternalServerError, err.Error()) return } w.WriteHeader(http.StatusNoContent) } // ---- Get-single and Delete handlers ---------------------------------------- func (s *Server) getFabric(w http.ResponseWriter, r *http.Request) { f, err := s.store.GetFabric(r.Context(), chi.URLParam(r, "name")) respondOne(w, f, err) } func (s *Server) deleteFabric(w http.ResponseWriter, r *http.Request) { respondDelete(w, s.store.DeleteFabric(r.Context(), chi.URLParam(r, "name"))) } func (s *Server) getZone(w http.ResponseWriter, r *http.Request) { z, err := s.store.GetZone(r.Context(), chi.URLParam(r, "name")) respondOne(w, z, err) } func (s *Server) deleteZone(w http.ResponseWriter, r *http.Request) { respondDelete(w, s.store.DeleteZone(r.Context(), chi.URLParam(r, "name"))) } func (s *Server) getAddressGroup(w http.ResponseWriter, r *http.Request) { g, err := s.store.GetAddressGroup(r.Context(), chi.URLParam(r, "name")) respondOne(w, g, err) } func (s *Server) deleteAddressGroup(w http.ResponseWriter, r *http.Request) { respondDelete(w, s.store.DeleteAddressGroup(r.Context(), chi.URLParam(r, "name"))) } func (s *Server) getPortGroup(w http.ResponseWriter, r *http.Request) { p, err := s.store.GetPortGroup(r.Context(), chi.URLParam(r, "name")) respondOne(w, p, err) } func (s *Server) deletePortGroup(w http.ResponseWriter, r *http.Request) { respondDelete(w, s.store.DeletePortGroup(r.Context(), chi.URLParam(r, "name"))) } func (s *Server) getDevice(w http.ResponseWriter, r *http.Request) { d, err := s.store.GetDevice(r.Context(), chi.URLParam(r, "name")) respondOne(w, d, err) } func (s *Server) deleteDevice(w http.ResponseWriter, r *http.Request) { respondDelete(w, s.store.DeleteDevice(r.Context(), chi.URLParam(r, "name"))) } func (s *Server) getBinding(w http.ResponseWriter, r *http.Request) { b, err := s.store.GetBinding(r.Context(), chi.URLParam(r, "name"), chi.URLParam(r, "zone")) respondOne(w, b, err) } func (s *Server) deleteBinding(w http.ResponseWriter, r *http.Request) { respondDelete(w, s.store.DeleteBinding(r.Context(), chi.URLParam(r, "name"), chi.URLParam(r, "zone"))) } func (s *Server) getRule(w http.ResponseWriter, r *http.Request) { id, err := strconv.ParseInt(chi.URLParam(r, "id"), 10, 64) if err != nil { writeError(w, http.StatusBadRequest, "id must be an integer") return } rule, err := s.store.GetRule(r.Context(), id) respondOne(w, rule, err) } // ---- Agent endpoints ------------------------------------------------------- func (s *Server) handleDeviceConfig(w http.ResponseWriter, r *http.Request) { cfg, err := compiler.Compile(r.Context(), s.store, chi.URLParam(r, "name")) if err != nil { if errors.Is(err, store.ErrNotFound) { writeError(w, http.StatusNotFound, "device not found") return } writeError(w, http.StatusInternalServerError, err.Error()) return } body, err := cfg.Marshal() if err != nil { writeError(w, http.StatusInternalServerError, err.Error()) return } w.Header().Set("Content-Type", "application/yaml") w.Header().Set("X-Tomswall-Generation", strconv.FormatInt(cfg.Generation, 10)) w.WriteHeader(http.StatusOK) _, _ = w.Write(body) } func (s *Server) handleDeviceStatus(w http.ResponseWriter, r *http.Request) { var body struct { Generation int64 `json:"generation"` } if !decode(w, r, &body) { return } if err := s.store.RecordDeviceStatus(r.Context(), chi.URLParam(r, "name"), body.Generation); err != nil { if errors.Is(err, store.ErrNotFound) { writeError(w, http.StatusNotFound, "device not found") return } writeError(w, http.StatusInternalServerError, err.Error()) return } w.WriteHeader(http.StatusNoContent) } // ---- helpers --------------------------------------------------------------- // decode reads a JSON request body into v, writing a 400 on failure. It returns // false when the caller should stop. func decode(w http.ResponseWriter, r *http.Request, v any) bool { dec := json.NewDecoder(r.Body) dec.DisallowUnknownFields() if err := dec.Decode(v); err != nil { writeError(w, http.StatusBadRequest, "invalid JSON: "+err.Error()) return false } return true } // respondList writes a list result or a 500, normalizing a nil slice to []. func respondList[T any](w http.ResponseWriter, list []T, err error) { if err != nil { writeError(w, http.StatusInternalServerError, err.Error()) return } if list == nil { list = []T{} } writeJSON(w, http.StatusOK, list) }