only allow admin access via unix socket
This commit is contained in:
parent
39cae64854
commit
13780f1c74
5 changed files with 73 additions and 16 deletions
|
|
@ -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)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
|
|
@ -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 = ¬FoundHandler{s}
|
s.mux.NotFoundHandler = ¬FoundHandler{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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue