diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..6411b36 --- /dev/null +++ b/.gitignore @@ -0,0 +1,4 @@ +/main +/main.exe +/rcw +/rcw.exe diff --git a/1shared.go b/1shared.go new file mode 100644 index 0000000..d783a36 --- /dev/null +++ b/1shared.go @@ -0,0 +1,9 @@ +package main + +import ( + "os" + "path/filepath" +) + +var binPath, _ = os.Executable() // store binary path +var socketPath = "/tmp/" + filepath.Base(binPath) + "-rcwd.sock" // store UNIX socket path diff --git a/2server.go b/2server.go new file mode 100644 index 0000000..5d50665 --- /dev/null +++ b/2server.go @@ -0,0 +1,133 @@ +package main + +import ( + "bytes" + "crypto/sha256" + "errors" + "fmt" + "io" + "log" + "net" + "net/rpc" + "os" + "syscall" + "time" +) + +var daemonHash []byte + +// RCWService provides an RPC method. +type RCWService struct{} + +// StartDaemon should be called to start an RPC server. +func StartDaemon() { + // store the hash of the daemon binary + daemonHash = getFileHash(binPath) + + // remove the socket file if it already exists + if _, err := os.Stat(socketPath); err == nil { + if err := os.Remove(socketPath); err != nil { + log.Fatalf("Failed to remove existing socket: %v", err) + } + } + + // register RCWService with the RPC package + if err := rpc.Register(&RCWService{}); err != nil { + log.Fatalf("Error registering RPC service: %v", err) + } + + // listen on the Unix domain socket + listener, err := net.Listen("unix", socketPath) + if err != nil { + log.Fatalf("Failed to listen on UNIX socket %s: %v", socketPath, err) + } + defer listener.Close() + log.Printf("RPC daemon listening on unix://%s", socketPath) + + // Accept connections (timeout after 3 minutes of inactivity) + for { + listener.(*net.UnixListener).SetDeadline(time.Now().Add(3 * time.Minute)) + + conn, err := listener.Accept() + if err != nil { + if err.(net.Error).Timeout() { + log.Println("Three minutes have passed without any connections. Exiting...") + os.Exit(0) + } + log.Printf("Accept error: %v", err) + continue + } + // use a goroutine to check the client's identity + go handleConn(conn) + } +} + +// GetPass is the RPC method. +// For now (as a test/example), it returns "hello" if the input is "hi". +func (h *RCWService) GetPass(request string, reply *string) error { + if request == "hi" { + *reply = "hello" + return nil + } + return errors.New("unexpected input, expected \"hi\"") +} + +// handleConn verifies the identity of the client. +// It uses the file descriptor of the connection to get the PID of the client, +// which is then used to get the path of the client's executable and calculate its hash. +// The passphrase is only returned if the client's executable hash matches the daemon's hash. +// This ensures that only the binary the daemon is embedded in can retrieve the passphrase. +func handleConn(conn net.Conn) { + // ensure access to the underlying file descriptor + unixConn, ok := conn.(*net.UnixConn) + if !ok { + log.Printf("Connection is not a Unix domain socket") + conn.Close() + return + } + // obtain SyscallConn for direct access to the file descriptor + rawConn, err := unixConn.SyscallConn() + if err != nil { + log.Printf("SyscallConn error: %v", err) + conn.Close() + return + } + + var callingPID int32 + _ = rawConn.Control(func(fd uintptr) { + // use syscall.GetsockoptUcred to fetch credentials + ucred, err := syscall.GetsockoptUcred(int(fd), syscall.SOL_SOCKET, syscall.SO_PEERCRED) + if err != nil { + log.Printf("Error getting peer credentials: %v", err) + return + } + callingPID = ucred.Pid + }) + + // check if the RPC call is coming from an identical binary + callingBinPath := pidToPath(callingPID) + if bytes.Equal(getFileHash(callingBinPath), daemonHash) { + // valid client; hand off the connection to the RPC server + rpc.ServeConn(conn) + } else { + // invalid client; close the connection w/o a response, + // log the client's path, and kill the daemon + conn.Close() + log.Printf("Request received from invalid client: %s", callingBinPath) // TODO log to file + os.Exit(2) + } +} + +// pidToPath returns the path of the executable that has the given PID. +func pidToPath(pid int32) string { + path, _ := os.Readlink(fmt.Sprintf("/proc/%d/exe", pid)) + return path +} + +// getFileHash returns the SHA256 hash of the file at the given path. +func getFileHash(path string) []byte { + file, _ := os.Open(path) + hash := sha256.New() + io.Copy(hash, file) + return hash.Sum(nil) +} diff --git a/3client.go b/3client.go new file mode 100644 index 0000000..eec3c5b --- /dev/null +++ b/3client.go @@ -0,0 +1,30 @@ +package main + +import ( + "log" + "net" + "net/rpc" +) + +func CallDaemon() string { + // connect to the UNIX domain socket + conn, err := net.Dial("unix", socketPath) + if err != nil { + log.Fatalf("Dial error: %v", err) + } + defer conn.Close() + + // create an RPC client using the connection + client := rpc.NewClient(conn) + defer client.Close() + + // request the passphrase from the RPC server + var reply string + err = client.Call("RCWService.GetPass", "hi", &reply) + if err != nil { + log.Fatalf("Error calling RCWService.GetPass: %v", err) + } + + // return the passphrase + return reply +} diff --git a/4mainTempExample.go b/4mainTempExample.go new file mode 100644 index 0000000..68811c9 --- /dev/null +++ b/4mainTempExample.go @@ -0,0 +1,19 @@ +package main + +import ( + "fmt" + "os" +) + +func main() { + if len(os.Args) > 1 { + switch os.Args[1] { + case "client": + fmt.Println("RPC Reply: " + CallDaemon()) + default: + StartDaemon() + } + } else { + StartDaemon() + } +} diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..94e1dff --- /dev/null +++ b/go.mod @@ -0,0 +1,3 @@ +module rcw + +go 1.24.1