only allow admin access via unix socket

This commit is contained in:
ston1th 2023-03-15 12:54:21 +01:00
commit 13780f1c74
5 changed files with 73 additions and 16 deletions

View file

@ -8,6 +8,7 @@ import (
"fmt" "fmt"
"os" "os"
"os/signal" "os/signal"
"path/filepath"
"syscall" "syscall"
"time" "time"
@ -23,6 +24,7 @@ var (
version string version string
listen string listen string
socket string
path string path string
log logr.Logger log logr.Logger
@ -58,6 +60,7 @@ func main() {
fs := flag.NewFlagSet("", flag.ExitOnError) fs := flag.NewFlagSet("", flag.ExitOnError)
klog.InitFlags(fs) klog.InitFlags(fs)
fs.StringVar(&listen, "listen", ":7070", "listen ip:port") fs.StringVar(&listen, "listen", ":7070", "listen ip:port")
fs.StringVar(&socket, "socket", "keyctl.sock", "listen unix domain socket")
fs.StringVar(&path, "path", "/var/keyctl", "db storage dir") fs.StringVar(&path, "path", "/var/keyctl", "db storage dir")
fs.Parse(os.Args[2:]) fs.Parse(os.Args[2:])
default: default:
@ -66,7 +69,7 @@ func main() {
log = klogr.New().WithName("main") log = klogr.New().WithName("main")
log.Info("starting keyctl", "version", version) log.Info("starting keyctl", "version", version)
srv, err := api.NewServer(klogr.New().WithName("srv"), listen) srv, err := api.NewServer(klogr.New().WithName("srv"), listen, filepath.Join(path, socket))
if err != nil { if err != nil {
klog.Fatalf("init failed: %s", err) klog.Fatalf("init failed: %s", err)
} }

View file

@ -11,6 +11,7 @@ import (
"net" "net"
"net/http" "net/http"
"net/http/httputil" "net/http/httputil"
"net/url"
"os" "os"
"strings" "strings"
"time" "time"
@ -19,6 +20,8 @@ import (
type Client struct { type Client struct {
endpoint string endpoint string
httpClient *http.Client httpClient *http.Client
dialer *net.Dialer
dialerFunc func(ctx context.Context, network, addr string) (net.Conn, error)
debugWriter io.Writer debugWriter io.Writer
//insecure bool //insecure bool
} }
@ -31,6 +34,22 @@ type ClientOption func(*Client)
func WithEndpoint(endpoint string) ClientOption { func WithEndpoint(endpoint string) ClientOption {
return func(client *Client) { return func(client *Client) {
u, err := url.Parse(endpoint)
if err != nil {
return
}
client.dialer = &net.Dialer{
Timeout: 30 * time.Second,
KeepAlive: 30 * time.Second,
}
if u.Scheme == "unix" {
client.dialerFunc = func(ctx context.Context, network, addr string) (net.Conn, error) {
return client.dialer.DialContext(ctx, "unix", u.Path)
}
client.endpoint = "http://unix"
return
}
client.dialerFunc = client.dialer.DialContext
client.endpoint = strings.TrimRight(endpoint, "/") client.endpoint = strings.TrimRight(endpoint, "/")
} }
} }
@ -104,10 +123,7 @@ func NewClient(options ...ClientOption) *Client {
return http.ErrUseLastResponse return http.ErrUseLastResponse
}, },
Transport: &http.Transport{ Transport: &http.Transport{
DialContext: (&net.Dialer{ DialContext: client.dialerFunc,
Timeout: 30 * time.Second,
KeepAlive: 30 * time.Second,
}).DialContext,
//ForceAttemptHTTP2: true, //ForceAttemptHTTP2: true,
MaxIdleConns: 100, MaxIdleConns: 100,
IdleConnTimeout: 90 * time.Second, IdleConnTimeout: 90 * time.Second,

View file

@ -6,6 +6,7 @@ import (
"context" "context"
"net" "net"
"net/http" "net/http"
"os"
"time" "time"
"git.giftfish.de/ston1th/keyctl/pkg/api/types" "git.giftfish.de/ston1th/keyctl/pkg/api/types"
@ -25,6 +26,7 @@ type Server struct {
Data *types.ContextData Data *types.ContextData
listen net.Listener listen net.Listener
socket net.Listener
} }
type notFoundHandler struct { type notFoundHandler struct {
@ -36,7 +38,7 @@ func (nf *notFoundHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
} }
// NewHTTPServer returns a new HTTPServer // NewHTTPServer returns a new HTTPServer
func NewServer(log logr.Logger, listen string) (*Server, error) { func NewServer(log logr.Logger, listen, socket string) (*Server, error) {
s := &Server{ s := &Server{
mux: mux.NewRouter(), mux: mux.NewRouter(),
log: log, log: log,
@ -48,6 +50,13 @@ func NewServer(log logr.Logger, listen string) (*Server, error) {
} }
s.listen = l s.listen = l
sock, err := unixListener(socket)
net.Listen("unix", socket)
if err != nil {
return nil, err
}
s.socket = sock
s.mux.NotFoundHandler = &notFoundHandler{s} s.mux.NotFoundHandler = &notFoundHandler{s}
for _, v := range serverv1.Routes { for _, v := range serverv1.Routes {
s.mux.HandleFunc(v.Path, s.contextWrapper(v.Handler)).Methods(v.Methods...) s.mux.HandleFunc(v.Path, s.contextWrapper(v.Handler)).Methods(v.Methods...)
@ -55,6 +64,15 @@ func NewServer(log logr.Logger, listen string) (*Server, error) {
return s, nil return s, nil
} }
func unixListener(socket string) (sock net.Listener, err error) {
sock, err = net.Listen("unix", socket)
if err != nil {
return
}
err = os.Chmod(socket, 0660)
return
}
func (s *Server) Start(db *db.DB) { func (s *Server) Start(db *db.DB) {
s.srv = &http.Server{ s.srv = &http.Server{
Handler: s.mux, Handler: s.mux,
@ -68,6 +86,12 @@ func (s *Server) Start(db *db.DB) {
s.log.Error(err, "") s.log.Error(err, "")
} }
}() }()
go func() {
err := s.srv.Serve(s.socket)
if err != nil && err != http.ErrServerClosed {
s.log.Error(err, "")
}
}()
return return
} }

View file

@ -40,6 +40,8 @@ type Context struct {
Data *ContextData Data *ContextData
Log logr.Logger Log logr.Logger
ip string
} }
// Method returns the request method // Method returns the request method
@ -69,7 +71,15 @@ func (c *Context) Form(name string) string {
} }
func (c *Context) ClientIP() (host string) { func (c *Context) ClientIP() (host string) {
host, _, _ = net.SplitHostPort(c.Request.RemoteAddr) host = c.Request.RemoteAddr
if host == "@" {
return
}
if c.ip != "" {
return c.ip
}
host, _, _ = net.SplitHostPort(host)
c.ip = host
return return
} }

View file

@ -5,7 +5,6 @@ package server
import ( import (
"encoding/json" "encoding/json"
"net/http" "net/http"
"net/netip"
"git.giftfish.de/ston1th/keyctl/pkg/api/types" "git.giftfish.de/ston1th/keyctl/pkg/api/types"
"git.giftfish.de/ston1th/keyctl/pkg/api/v1/schema" "git.giftfish.de/ston1th/keyctl/pkg/api/v1/schema"
@ -14,16 +13,21 @@ import (
func local(h types.CtxHandler) types.CtxHandler { func local(h types.CtxHandler) types.CtxHandler {
return func(ctx *types.Context) { return func(ctx *types.Context) {
addr, err := netip.ParseAddr(ctx.ClientIP()) ip := ctx.ClientIP()
if err != nil { if ip == "@" {
ctx.Err(types.ErrForbidden) h(ctx)
return return
} }
if !addr.IsLoopback() { //addr, err := netip.ParseAddr(ip)
ctx.Err(types.ErrForbidden) //if err != nil {
return // ctx.Err(types.ErrForbidden)
} // return
h(ctx) //}
//if addr.IsLoopback() {
// h(ctx)
// return
//}
ctx.Err(types.ErrForbidden)
} }
} }