diff --git a/.gitignore b/.gitignore index 3b3410c..a19a38b 100644 --- a/.gitignore +++ b/.gitignore @@ -4,3 +4,5 @@ _tmp/ .cache *.json *.out +*.html +*.log diff --git a/cmd/oneauth/commands/cmd_agent.go b/cmd/oneauth/commands/cmd_agent.go index 9adbe69..7ab7a67 100644 --- a/cmd/oneauth/commands/cmd_agent.go +++ b/cmd/oneauth/commands/cmd_agent.go @@ -24,14 +24,7 @@ var agentCmd = &cli.Command{ Usage: "SSH Agent", Description: "All configuration options can be set in the config file", Before: func(_ *cli.Context) error { - version := buildinfo.Version - - commit := buildinfo.Commit - if len(commit) > 8 { - version += "-" + commit[:8] - } - - log.Printf("OneAuth version: %s", version) + log.Printf("OneAuth version: %s", buildinfo.FormattedVersion()) return nil }, @@ -45,14 +38,26 @@ var agentCmd = &cli.Command{ log := logger.New(config.AgentLogPath) - if config.Keyring.Yubikey.Serial == 0 { - return fmt.Errorf("yubikey serial is required") - } - var agent *sshagent.SSHAgent switch config.Socket.Type { case "unix": + // Determine YubiKey serial to use + var yubikeySerial uint32 + if config.Keyring.DisableYubikey { + log.Println("YubiKey keyring disabled - running without hardware authentication") + yubikeySerial = 0 // Disabled mode + + } else { + // YubiKey serial is required for unix socket type when not disabled + if config.Keyring.Yubikey.Serial == 0 { + return fmt.Errorf("yubikey serial is required for unix socket type (or set disable_yubikey: true)") + } + yubikeySerial = config.Keyring.Yubikey.Serial + log.WithField("yubikey", yubikeySerial).Println("opening yubikey:", yubikeySerial) + } + + // Set up socket if _, err := os.Stat(config.Socket.Path); err == nil { os.Remove(config.Socket.Path) } @@ -61,9 +66,8 @@ var agentCmd = &cli.Command{ return fmt.Errorf("failed to create directory: %w", err) } - log.WithField("yubikey", config.Keyring.Yubikey.Serial).Println("opening yubikey:", config.Keyring.Yubikey.Serial) - - agent, err = sshagent.New(config.Keyring.Yubikey.Serial, log, config) + // Create agent + agent, err = sshagent.New(yubikeySerial, log, config) if err != nil { return fmt.Errorf("failed to create agent: %w", err) } @@ -73,7 +77,8 @@ var agentCmd = &cli.Command{ }) case "dummy": - log.Println("skipping socket creation") + log.Println("skipping socket creation - running in dummy mode") + // For dummy mode, we don't need YubiKey, so agent remains nil default: return fmt.Errorf("socket type %s is not supported", config.Socket.Type) diff --git a/cmd/oneauth/commands/cmd_agent_test.go b/cmd/oneauth/commands/cmd_agent_test.go new file mode 100644 index 0000000..ae23e5c --- /dev/null +++ b/cmd/oneauth/commands/cmd_agent_test.go @@ -0,0 +1,204 @@ +package commands + +import ( + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/vitalvas/oneauth/cmd/oneauth/config" +) + +func TestAgentValidation_YubikeySerialRequirement(t *testing.T) { + t.Run("UnixSocketRequiresYubikeySerial", func(t *testing.T) { + // Create a temporary config file with serial = 0 + tempDir := t.TempDir() + configPath := filepath.Join(tempDir, "config.yaml") + + configContent := ` +socket: + type: unix + path: ` + filepath.Join(tempDir, "test.sock") + ` +control_socket_path: ` + filepath.Join(tempDir, "test_ctrl.sock") + ` +agent_log_path: ` + filepath.Join(tempDir, "test.log") + ` +keyring: + yubikey: + serial: 0 +` + err := os.WriteFile(configPath, []byte(configContent), 0644) + assert.NoError(t, err) + + // Load the config + cfg, err := config.Load(configPath) + assert.NoError(t, err) + + // Test the validation logic that's in the agent command + if cfg.Socket.Type == "unix" && !cfg.Keyring.DisableYubikey && cfg.Keyring.Yubikey.Serial == 0 { + err = assert.AnError // Simulate the error that would be returned + } + + // Should detect that YubiKey serial is required + assert.Error(t, err, "Expected error when YubiKey serial is 0 for unix socket") + }) + + t.Run("DummySocketDoesNotRequireYubikeySerial", func(t *testing.T) { + // Create a temporary config file with dummy socket and serial = 0 + tempDir := t.TempDir() + configPath := filepath.Join(tempDir, "config.yaml") + + configContent := ` +socket: + type: dummy +control_socket_path: ` + filepath.Join(tempDir, "test_ctrl.sock") + ` +agent_log_path: ` + filepath.Join(tempDir, "test.log") + ` +keyring: + yubikey: + serial: 0 +` + err := os.WriteFile(configPath, []byte(configContent), 0644) + assert.NoError(t, err) + + // Load the config + cfg, err := config.Load(configPath) + assert.NoError(t, err) + + // Test the validation logic - dummy socket should not require YubiKey + var validationErr error + if cfg.Socket.Type == "unix" && !cfg.Keyring.DisableYubikey && cfg.Keyring.Yubikey.Serial == 0 { + validationErr = assert.AnError // This should NOT trigger for dummy + } + + // Should NOT require YubiKey serial for dummy socket + assert.NoError(t, validationErr, "Dummy socket should not require YubiKey serial") + assert.Equal(t, "dummy", cfg.Socket.Type) + assert.Equal(t, uint32(0), cfg.Keyring.Yubikey.Serial) + }) + + t.Run("UnixSocketWithValidSerial", func(t *testing.T) { + // Create a temporary config file with valid serial + tempDir := t.TempDir() + configPath := filepath.Join(tempDir, "config.yaml") + + configContent := ` +socket: + type: unix + path: ` + filepath.Join(tempDir, "test.sock") + ` +control_socket_path: ` + filepath.Join(tempDir, "test_ctrl.sock") + ` +agent_log_path: ` + filepath.Join(tempDir, "test.log") + ` +keyring: + yubikey: + serial: 12345 +` + err := os.WriteFile(configPath, []byte(configContent), 0644) + assert.NoError(t, err) + + // Load the config + cfg, err := config.Load(configPath) + assert.NoError(t, err) + + // Test the validation logic + var validationErr error + if cfg.Socket.Type == "unix" && !cfg.Keyring.DisableYubikey && cfg.Keyring.Yubikey.Serial == 0 { + validationErr = assert.AnError // This should NOT trigger with valid serial + } + + // Should pass validation with valid serial + assert.NoError(t, validationErr) + assert.Equal(t, "unix", cfg.Socket.Type) + assert.Equal(t, uint32(12345), cfg.Keyring.Yubikey.Serial) + }) + + t.Run("UnixSocketWithDisabledYubikey", func(t *testing.T) { + // Create a temporary config file with disabled YubiKey + tempDir := t.TempDir() + configPath := filepath.Join(tempDir, "config.yaml") + + configContent := ` +socket: + type: unix + path: ` + filepath.Join(tempDir, "test.sock") + ` +control_socket_path: ` + filepath.Join(tempDir, "test_ctrl.sock") + ` +agent_log_path: ` + filepath.Join(tempDir, "test.log") + ` +keyring: + disable_yubikey: true + yubikey: + serial: 0 +` + err := os.WriteFile(configPath, []byte(configContent), 0644) + assert.NoError(t, err) + + // Load the config + cfg, err := config.Load(configPath) + assert.NoError(t, err) + + // Test the validation logic - should not require serial when disabled + var validationErr error + if cfg.Socket.Type == "unix" && !cfg.Keyring.DisableYubikey && cfg.Keyring.Yubikey.Serial == 0 { + validationErr = assert.AnError // This should NOT trigger when disabled + } + + // Should pass validation when YubiKey is disabled + assert.NoError(t, validationErr) + assert.Equal(t, "unix", cfg.Socket.Type) + assert.True(t, cfg.Keyring.DisableYubikey) + assert.Equal(t, uint32(0), cfg.Keyring.Yubikey.Serial) + }) + + t.Run("UnixSocketWithDisabledYubikeyButValidSerial", func(t *testing.T) { + // Create a temporary config file with disabled YubiKey but valid serial + tempDir := t.TempDir() + configPath := filepath.Join(tempDir, "config.yaml") + + configContent := ` +socket: + type: unix + path: ` + filepath.Join(tempDir, "test.sock") + ` +control_socket_path: ` + filepath.Join(tempDir, "test_ctrl.sock") + ` +agent_log_path: ` + filepath.Join(tempDir, "test.log") + ` +keyring: + disable_yubikey: true + yubikey: + serial: 12345 +` + err := os.WriteFile(configPath, []byte(configContent), 0644) + assert.NoError(t, err) + + // Load the config + cfg, err := config.Load(configPath) + assert.NoError(t, err) + + // Should pass validation - when disabled, serial is ignored + assert.Equal(t, "unix", cfg.Socket.Type) + assert.True(t, cfg.Keyring.DisableYubikey) + assert.Equal(t, uint32(12345), cfg.Keyring.Yubikey.Serial) // Serial is present but ignored + }) +} + +func TestAgentValidation_SocketTypes(t *testing.T) { + t.Run("ValidSocketTypes", func(t *testing.T) { + validTypes := []string{"unix", "dummy"} + + for _, socketType := range validTypes { + tempDir := t.TempDir() + configPath := filepath.Join(tempDir, "config.yaml") + + configContent := ` +socket: + type: ` + socketType + ` + path: ` + filepath.Join(tempDir, "test.sock") + ` +control_socket_path: ` + filepath.Join(tempDir, "test_ctrl.sock") + ` +agent_log_path: ` + filepath.Join(tempDir, "test.log") + ` +keyring: + yubikey: + serial: 12345 +` + err := os.WriteFile(configPath, []byte(configContent), 0644) + assert.NoError(t, err) + + // Load the config + cfg, err := config.Load(configPath) + assert.NoError(t, err) + assert.Equal(t, socketType, cfg.Socket.Type) + } + }) +} \ No newline at end of file diff --git a/cmd/oneauth/commands/cmd_info.go b/cmd/oneauth/commands/cmd_info.go index 2f5c5ff..69dd12d 100644 --- a/cmd/oneauth/commands/cmd_info.go +++ b/cmd/oneauth/commands/cmd_info.go @@ -7,6 +7,8 @@ import ( "strings" "github.com/urfave/cli/v2" + "github.com/vitalvas/oneauth/cmd/oneauth/config" + "github.com/vitalvas/oneauth/cmd/oneauth/rpclient" "github.com/vitalvas/oneauth/internal/yubikey" ) @@ -15,6 +17,13 @@ const infoTmpl = ` {{- range $key := .Keys }} - {{ $key.Name }} (Serial: {{ $key.Serial }}, Version: {{ $key.Version }}) {{- end }} +{{- if .AgentPid }} + +--- Agent --- +Agent PID: {{ .AgentPid }} +Version: {{ .AgentVersion }} +Uptime: {{ .AgentUptime }} +{{- end }} ` type InfoKey struct { @@ -24,20 +33,71 @@ type InfoKey struct { } type infoData struct { - Keys []InfoKey + Keys []InfoKey + AgentPid int + AgentVersion string + AgentUptime string } var infoCmd = &cli.Command{ Name: "info", Usage: "Prints detailed information", - Action: func(_ *cli.Context) error { + Before: loadConfig, + Action: func(c *cli.Context) error { + info := infoData{} + + // Try to get info from running agent via control socket first + if globalConfig != nil { + client, err := rpclient.New(globalConfig.ControlSocketPath) + if err == nil { + defer client.Close() + + rpcInfo, err := client.GetInfo() + if err == nil { + info.AgentPid = rpcInfo.Pid + info.AgentVersion = rpcInfo.Version + info.AgentUptime = rpcInfo.Uptime + // Use keys from the running agent + for _, key := range rpcInfo.Keys { + info.Keys = append(info.Keys, InfoKey{ + Name: key.Name, + Serial: key.Serial, + Version: key.Version, + }) + } + + render := func(tmpl string, data interface{}) string { + var out strings.Builder + if err := template.Must(template.New("tmpl").Parse(tmpl)).Execute(&out, data); err != nil { + log.Fatalf("failed to render template: %v", err) + } + return strings.TrimSpace(out.String()) + } + + fmt.Println(render(infoTmpl, info)) + return nil + } + } + } + + // Fallback to local YubiKey detection when control socket is unavailable + var configPath string + if globalConfig == nil { + configPath = c.String("config") + cfg, err := config.Load(configPath) + if err != nil { + // Continue with YubiKey detection even if config fails + log.Printf("Warning: failed to load config: %v", err) + } else { + globalConfig = cfg + } + } + cards, err := yubikey.Cards() if err != nil { return err } - info := infoData{} - for _, card := range cards { info.Keys = append(info.Keys, InfoKey{ Name: card.String(), @@ -48,16 +108,13 @@ var infoCmd = &cli.Command{ render := func(tmpl string, data interface{}) string { var out strings.Builder - if err := template.Must(template.New("tmpl").Parse(tmpl)).Execute(&out, data); err != nil { log.Fatalf("failed to render template: %v", err) } - return strings.TrimSpace(out.String()) } fmt.Println(render(infoTmpl, info)) - return nil }, } diff --git a/cmd/oneauth/commands/cmd_service_enable.go b/cmd/oneauth/commands/cmd_service_enable.go index 559f0d9..6df7948 100644 --- a/cmd/oneauth/commands/cmd_service_enable.go +++ b/cmd/oneauth/commands/cmd_service_enable.go @@ -1,7 +1,9 @@ package commands import ( + "errors" "fmt" + "time" "github.com/urfave/cli/v2" "github.com/vitalvas/oneauth/cmd/oneauth/service" @@ -15,8 +17,16 @@ var serviceEnableCmd = &cli.Command{ return err } - fmt.Println("oneauth agent service has been successfully installed and probably started") + for i := 0; i < 100; i++ { + if service.IsRunning() { + fmt.Println("oneauth agent service has been successfully started") + return nil + } - return nil + fmt.Println("waiting for oneauth agent service to start...") + time.Sleep(500 * time.Millisecond) + } + + return errors.New("oneauth agent service has been installed, but it is not running") }, } diff --git a/cmd/oneauth/commands/cmd_service_restart.go b/cmd/oneauth/commands/cmd_service_restart.go index f97854e..e29bd58 100644 --- a/cmd/oneauth/commands/cmd_service_restart.go +++ b/cmd/oneauth/commands/cmd_service_restart.go @@ -1,7 +1,9 @@ package commands import ( + "errors" "fmt" + "time" "github.com/urfave/cli/v2" "github.com/vitalvas/oneauth/cmd/oneauth/service" @@ -15,8 +17,16 @@ var serviceRestartCmd = &cli.Command{ return err } - fmt.Println("done...") + for i := 0; i < 100; i++ { + if service.IsRunning() { + fmt.Println("oneauth agent service has been successfully restarted") + return nil + } - return nil + fmt.Println("waiting for oneauth agent service to start...") + time.Sleep(500 * time.Millisecond) + } + + return errors.New("oneauth agent service failed to restart") }, } diff --git a/cmd/oneauth/commands/tools_config.go b/cmd/oneauth/commands/tools_config.go new file mode 100644 index 0000000..24f3f58 --- /dev/null +++ b/cmd/oneauth/commands/tools_config.go @@ -0,0 +1,28 @@ +package commands + +import ( + "os" + + "github.com/urfave/cli/v2" + "github.com/vitalvas/oneauth/cmd/oneauth/config" +) + +var globalConfig *config.Config + +func loadConfig(c *cli.Context) error { + configPath := c.String("config") + + if _, err := os.Stat(configPath); err != nil { + if os.IsNotExist(err) { + return nil + } + + return err + } + + var err error + + globalConfig, err = config.Load(configPath) + + return err +} diff --git a/cmd/oneauth/commands/tools_config_test.go b/cmd/oneauth/commands/tools_config_test.go new file mode 100644 index 0000000..ab68503 --- /dev/null +++ b/cmd/oneauth/commands/tools_config_test.go @@ -0,0 +1,48 @@ +package commands + +import ( + "flag" + "os" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/urfave/cli/v2" +) + +func TestLoadConfig(t *testing.T) { + t.Run("FileDoesNotExist", func(t *testing.T) { + app := cli.NewApp() + set := flag.NewFlagSet("test", 0) + set.String("config", "/non/existent/path", "doc") + c := cli.NewContext(app, set, nil) + + err := loadConfig(c) + assert.Nil(t, err) + assert.Nil(t, globalConfig) + }) + + t.Run("FileExists", func(t *testing.T) { + tmpFile, err := os.CreateTemp("", "config-*.yml") + if err != nil { + assert.Error(t, err) + return + } + defer os.Remove(tmpFile.Name()) + + _, err = tmpFile.WriteString(`control_socket_path: /tmp/oneauth.sock`) + if err != nil { + assert.Error(t, err) + return + } + tmpFile.Close() + + app := cli.NewApp() + set := flag.NewFlagSet("test", 0) + set.String("config", tmpFile.Name(), "doc") + c := cli.NewContext(app, set, nil) + + err = loadConfig(c) + assert.Nil(t, err) + assert.NotNil(t, globalConfig) + }) +} diff --git a/cmd/oneauth/config/config.go b/cmd/oneauth/config/config.go index 0220f3e..bdca4d5 100644 --- a/cmd/oneauth/config/config.go +++ b/cmd/oneauth/config/config.go @@ -6,10 +6,20 @@ import ( "time" "github.com/vitalvas/oneauth/cmd/oneauth/paths" + "github.com/vitalvas/oneauth/internal/tools" "gopkg.in/yaml.v3" ) func Load(filePath string) (*Config, error) { + rootDir, err := paths.RootDir() + if err != nil { + return nil, err + } + + if err := tools.MkDir(rootDir, 0700); err != nil { + return nil, err + } + agentSocketPath, err := paths.AgentSocket() if err != nil { return nil, err @@ -61,3 +71,4 @@ func loadYamlFile(filePath string, v *Config) error { return decoder.Decode(v) } + diff --git a/cmd/oneauth/config/config_test.go b/cmd/oneauth/config/config_test.go index afbf05b..6333926 100644 --- a/cmd/oneauth/config/config_test.go +++ b/cmd/oneauth/config/config_test.go @@ -447,3 +447,85 @@ func TestLoadYamlFile_EdgeCases(t *testing.T) { assert.Error(t, err) }) } + +// setupTestHomeDir creates a temporary home directory for testing +func setupTestHomeDir(t *testing.T) (string, func()) { + tmpHome, err := os.MkdirTemp("", "test-home-*") + require.NoError(t, err) + + // Create .oneauth directory in temp home + oneauthDir := filepath.Join(tmpHome, ".oneauth") + err = os.MkdirAll(oneauthDir, 0755) + require.NoError(t, err) + + // Temporarily set HOME env var + originalHome := os.Getenv("HOME") + os.Setenv("HOME", tmpHome) + + cleanup := func() { + os.Setenv("HOME", originalHome) + os.RemoveAll(tmpHome) + } + + return tmpHome, cleanup +} + +func TestControlSocketPathGeneration(t *testing.T) { + t.Run("AutoGeneration", func(t *testing.T) { + _, cleanup := setupTestHomeDir(t) + defer cleanup() + + // Create config with custom socket path but no control_socket_path + tmpFile, err := os.CreateTemp("", "config-*.yaml") + require.NoError(t, err) + defer os.Remove(tmpFile.Name()) + defer tmpFile.Close() + + configContent := ` +socket: + type: "unix" + path: "/tmp/custom-agent.sock" +keyring: + disable_yubikey: true +` + _, err = tmpFile.Write([]byte(configContent)) + require.NoError(t, err) + + config, err := Load(tmpFile.Name()) + require.NoError(t, err) + + // Should use default control socket path when not explicitly set + assert.Contains(t, config.ControlSocketPath, "oneauth-ctrl.sock") + assert.Equal(t, "/tmp/custom-agent.sock", config.Socket.Path) + }) + + t.Run("ExplicitControlSocketPath", func(t *testing.T) { + _, cleanup := setupTestHomeDir(t) + defer cleanup() + + // Create config with both socket path and explicit control_socket_path + tmpFile, err := os.CreateTemp("", "config-*.yaml") + require.NoError(t, err) + defer os.Remove(tmpFile.Name()) + defer tmpFile.Close() + + configContent := ` +control_socket_path: "/tmp/explicit-control.sock" +socket: + type: "unix" + path: "/tmp/custom-agent.sock" +keyring: + disable_yubikey: true +` + _, err = tmpFile.Write([]byte(configContent)) + require.NoError(t, err) + + config, err := Load(tmpFile.Name()) + require.NoError(t, err) + + // Should use explicit control socket path + assert.Equal(t, "/tmp/explicit-control.sock", config.ControlSocketPath) + assert.Equal(t, "/tmp/custom-agent.sock", config.Socket.Path) + }) +} + diff --git a/cmd/oneauth/config/struct.go b/cmd/oneauth/config/struct.go index 31bc21a..00199cb 100644 --- a/cmd/oneauth/config/struct.go +++ b/cmd/oneauth/config/struct.go @@ -17,6 +17,7 @@ type Socket struct { } type Keyring struct { + DisableYubikey bool `yaml:"disable_yubikey,omitempty"` Yubikey KeyringYubikey `yaml:"yubikey,omitempty"` BeforeSignHook string `yaml:"before_sign_hook,omitempty"` KeepKeySeconds int64 `yaml:"keep_key_seconds,omitempty"` diff --git a/cmd/oneauth/integration/rpc_integration_test.go b/cmd/oneauth/integration/rpc_integration_test.go new file mode 100644 index 0000000..4a71e34 --- /dev/null +++ b/cmd/oneauth/integration/rpc_integration_test.go @@ -0,0 +1,140 @@ +package integration + +import ( + "context" + "os" + "path/filepath" + "testing" + "time" + + "github.com/sirupsen/logrus" + "github.com/stretchr/testify/assert" + "github.com/vitalvas/oneauth/cmd/oneauth/rpclient" + "github.com/vitalvas/oneauth/cmd/oneauth/rpcserver" + "github.com/vitalvas/oneauth/cmd/oneauth/sshagent" +) + +func TestRPCIntegration(t *testing.T) { + t.Run("FullServerClientIntegration", func(t *testing.T) { + // Create temporary directory + tempDir, err := os.MkdirTemp("", "integration_test") + assert.NoError(t, err) + defer os.RemoveAll(tempDir) + + socketPath := filepath.Join(tempDir, "test.sock") + + // Create RPC server + sshAgent := &sshagent.SSHAgent{} + log := logrus.New() + log.SetLevel(logrus.FatalLevel) // Reduce log noise in tests + rpcServer := rpcserver.New(sshAgent, log) + + // Start server in goroutine + serverErrChan := make(chan error) + go func() { + serverErrChan <- rpcServer.ListenAndServe(context.Background(), socketPath) + }() + + // Give server time to start + time.Sleep(100 * time.Millisecond) + + // Verify socket was created + _, err = os.Stat(socketPath) + assert.NoError(t, err) + + // Create client and test connection + client, err := rpclient.New(socketPath) + assert.NoError(t, err) + assert.NotNil(t, client) + defer client.Close() + + // Test RPC call + info, err := client.GetInfo() + assert.NoError(t, err) + assert.NotNil(t, info) + assert.Equal(t, os.Getpid(), info.Pid) + + // Test multiple concurrent calls + doneChan := make(chan bool, 5) + for i := 0; i < 5; i++ { + go func() { + info, err := client.GetInfo() + assert.NoError(t, err) + assert.NotNil(t, info) + assert.Equal(t, os.Getpid(), info.Pid) + doneChan <- true + }() + } + + // Wait for all goroutines to complete + for i := 0; i < 5; i++ { + <-doneChan + } + + // Close client + err = client.Close() + assert.NoError(t, err) + + // Shutdown server + rpcServer.Shutdown() + + // Wait for server to shutdown + select { + case err := <-serverErrChan: + // Should complete with listener closed error (expected) + if err != nil { + assert.Contains(t, err.Error(), "use of closed network connection") + } + case <-time.After(5 * time.Second): + t.Fatal("Server did not shutdown in time") + } + }) + + t.Run("ClientConnectToNonExistentServer", func(t *testing.T) { + nonExistentSocketPath := "/tmp/non_existent_socket.sock" + + client, err := rpclient.New(nonExistentSocketPath) + assert.Error(t, err) + assert.Nil(t, client) + }) + + t.Run("ClientCallAfterServerShutdown", func(t *testing.T) { + // Create temporary directory + tempDir, err := os.MkdirTemp("", "integration_test") + assert.NoError(t, err) + defer os.RemoveAll(tempDir) + + socketPath := filepath.Join(tempDir, "test.sock") + + // Create and start server + log := logrus.New() + log.SetLevel(logrus.FatalLevel) + server := rpcserver.New(nil, log) + + go func() { + server.ListenAndServe(context.Background(), socketPath) + }() + + // Give server time to start + time.Sleep(100 * time.Millisecond) + + // Create client and verify it works + client, err := rpclient.New(socketPath) + assert.NoError(t, err) + + info, err := client.GetInfo() + assert.NoError(t, err) + assert.NotNil(t, info) + + client.Close() + + // Shutdown server + server.Shutdown() + time.Sleep(100 * time.Millisecond) + + // Try to create new client after shutdown - should fail + client2, err := rpclient.New(socketPath) + assert.Error(t, err) + assert.Nil(t, client2) + }) +} diff --git a/cmd/oneauth/paths/paths.go b/cmd/oneauth/paths/paths.go index 4e1928e..c6bf7b3 100644 --- a/cmd/oneauth/paths/paths.go +++ b/cmd/oneauth/paths/paths.go @@ -13,7 +13,7 @@ func AgentSocket() (string, error) { } func ControlSocket() (string, error) { - return tools.InHomeDir(oneauthDir, "control.sock") + return tools.InHomeDir(oneauthDir, "oneauth-ctrl.sock") } func Config() (string, error) { @@ -27,3 +27,7 @@ func BinDir() (string, error) { func LogDir() (string, error) { return tools.InHomeDir(oneauthDir, "log") } + +func RootDir() (string, error) { + return tools.InHomeDir(oneauthDir) +} diff --git a/cmd/oneauth/paths/paths_test.go b/cmd/oneauth/paths/paths_test.go index ca70feb..3798111 100644 --- a/cmd/oneauth/paths/paths_test.go +++ b/cmd/oneauth/paths/paths_test.go @@ -39,7 +39,7 @@ func TestControlSocket(t *testing.T) { actual, err := ControlSocket() assert.Nil(t, err, "Error getting control socket path: %v", err) - expected := filepath.Join(home, oneauthDir, "control.sock") + expected := filepath.Join(home, oneauthDir, "oneauth-ctrl.sock") assert.Equal(t, expected, actual, "Expected result: %s, got result: %s", expected, actual) } @@ -102,3 +102,15 @@ func TestServiceFile(t *testing.T) { t.Errorf("Expected result to be in %s, got result: %s", correctDir, path) } } + +func TestRootDir(t *testing.T) { + home, err := os.UserHomeDir() + assert.Nil(t, err, "Error getting user home directory: %v", err) + + actual, err := RootDir() + assert.Nil(t, err, "Error getting root directory: %v", err) + + expected := filepath.Join(home, oneauthDir) + + assert.Equal(t, expected, actual, "Expected result: %s, got result: %s", expected, actual) +} diff --git a/cmd/oneauth/rpclient/client.go b/cmd/oneauth/rpclient/client.go new file mode 100644 index 0000000..5bcb6b7 --- /dev/null +++ b/cmd/oneauth/rpclient/client.go @@ -0,0 +1,45 @@ +package rpclient + +import ( + "net" + "net/rpc" + "net/rpc/jsonrpc" + "os" +) + +type Client struct { + socketPath string + client *rpc.Client +} + +func New(socketPath string) (*Client, error) { + conn, err := net.Dial("unix", socketPath) + if err != nil { + return nil, err + } + + client := rpc.NewClientWithCodec(jsonrpc.NewClientCodec(conn)) + + return &Client{ + socketPath: socketPath, + client: client, + }, nil +} + +func (c *Client) Close() error { + if c.client != nil { + return c.client.Close() + } + return nil +} + +func IsClient(socketPath string) bool { + if info, err := os.Stat(socketPath); err != nil { + return false + } else if info.Mode()&os.ModeSocket == 0 { + return false + } + + _, err := net.Dial("unix", socketPath) + return err == nil +} diff --git a/cmd/oneauth/rpclient/client_test.go b/cmd/oneauth/rpclient/client_test.go new file mode 100644 index 0000000..844417a --- /dev/null +++ b/cmd/oneauth/rpclient/client_test.go @@ -0,0 +1,105 @@ +package rpclient + +import ( + "net" + "net/rpc" + "net/rpc/jsonrpc" + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestNew(t *testing.T) { + t.Run("ValidSocket", func(t *testing.T) { + validSocketPath := filepath.Join(t.TempDir(), "valid.sock") + listener, err := net.Listen("unix", validSocketPath) + if err != nil { + assert.NoError(t, err) + return + } + defer os.Remove(validSocketPath) + defer listener.Close() + + // Start a simple RPC server + go func() { + for { + conn, err := listener.Accept() + if err != nil { + return + } + rpc.ServeCodec(jsonrpc.NewServerCodec(conn)) + } + }() + + client, err := New(validSocketPath) + assert.NoError(t, err) + assert.NotNil(t, client) + if client != nil { + client.Close() + } + }) + + t.Run("InvalidSocket", func(t *testing.T) { + invalidSocketPath := filepath.Join(t.TempDir(), "invalid.sock") + file, err := os.Create(invalidSocketPath) + if err != nil { + assert.NoError(t, err) + return + } + defer os.Remove(invalidSocketPath) + defer file.Close() + + client, err := New(invalidSocketPath) + assert.Error(t, err) + assert.Nil(t, client) + }) + + t.Run("NonExistentSocket", func(t *testing.T) { + nonExistentSocketPath := filepath.Join(t.TempDir(), "non_existent_socket.sock") + + client, err := New(nonExistentSocketPath) + assert.Error(t, err) + assert.Nil(t, client) + }) +} + +func TestIsClient(t *testing.T) { + t.Run("ValidSocket", func(t *testing.T) { + validSocketPath := filepath.Join(t.TempDir(), "valid.sock") + listener, err := net.Listen("unix", validSocketPath) + if err != nil { + assert.NoError(t, err) + return + } + defer os.Remove(validSocketPath) + defer listener.Close() + + assert.True(t, IsClient(validSocketPath)) + }) + + t.Run("InvalidSocket", func(t *testing.T) { + invalidSocketPath := filepath.Join(t.TempDir(), "invalid.sock") + file, err := os.Create(invalidSocketPath) + if err != nil { + assert.NoError(t, err) + return + } + defer os.Remove(invalidSocketPath) + defer file.Close() + + assert.False(t, IsClient(invalidSocketPath)) + }) + + t.Run("NonExistentSocket", func(t *testing.T) { + nonExistentSocketPath := filepath.Join(t.TempDir(), "non_existent_socket.sock") + assert.False(t, IsClient(nonExistentSocketPath)) + }) + + t.Run("CloseClient", func(t *testing.T) { + client := &Client{} + err := client.Close() + assert.NoError(t, err) + }) +} diff --git a/cmd/oneauth/rpclient/rpc_info.go b/cmd/oneauth/rpclient/rpc_info.go new file mode 100644 index 0000000..37e7684 --- /dev/null +++ b/cmd/oneauth/rpclient/rpc_info.go @@ -0,0 +1,17 @@ +package rpclient + +import ( + "github.com/vitalvas/oneauth/cmd/oneauth/rpcserver" +) + +func (c *Client) GetInfo() (*rpcserver.InfoReply, error) { + args := &rpcserver.InfoArgs{} + reply := &rpcserver.InfoReply{} + + err := c.client.Call("AgentService.Info", args, reply) + if err != nil { + return nil, err + } + + return reply, nil +} diff --git a/cmd/oneauth/rpclient/rpc_info_test.go b/cmd/oneauth/rpclient/rpc_info_test.go new file mode 100644 index 0000000..8594914 --- /dev/null +++ b/cmd/oneauth/rpclient/rpc_info_test.go @@ -0,0 +1,80 @@ +package rpclient + +import ( + "net" + "net/rpc" + "net/rpc/jsonrpc" + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/vitalvas/oneauth/cmd/oneauth/rpcserver" +) + +type MockAgentService struct{} + +func (m *MockAgentService) Info(_ *rpcserver.InfoArgs, reply *rpcserver.InfoReply) error { + reply.Pid = 12345 + reply.Version = "1.0.0-abcd1234" + reply.Uptime = "1h30m0s" + reply.Keys = []rpcserver.InfoKey{ + { + Name: "YubiKey 5 NFC", + Serial: "12345", + Version: "5.4.3", + }, + } + return nil +} + +func TestGetInfo(t *testing.T) { + t.Run("successful response", func(t *testing.T) { + // Create temporary socket + socketPath := filepath.Join(t.TempDir(), "test.sock") + listener, err := net.Listen("unix", socketPath) + assert.NoError(t, err) + defer os.Remove(socketPath) + defer listener.Close() + + // Start RPC server + rpcServer := rpc.NewServer() + rpcServer.RegisterName("AgentService", &MockAgentService{}) + + go func() { + for { + conn, err := listener.Accept() + if err != nil { + return + } + go rpcServer.ServeCodec(jsonrpc.NewServerCodec(conn)) + } + }() + + // Create client + client, err := New(socketPath) + assert.NoError(t, err) + assert.NotNil(t, client) + defer client.Close() + + // Test GetInfo + info, err := client.GetInfo() + assert.NoError(t, err) + assert.NotNil(t, info) + assert.Equal(t, 12345, info.Pid) + assert.Equal(t, "1.0.0-abcd1234", info.Version) + assert.Equal(t, "1h30m0s", info.Uptime) + assert.Len(t, info.Keys, 1) + assert.Equal(t, "YubiKey 5 NFC", info.Keys[0].Name) + assert.Equal(t, "12345", info.Keys[0].Serial) + assert.Equal(t, "5.4.3", info.Keys[0].Version) + }) + + t.Run("connection error", func(t *testing.T) { + nonExistentSocketPath := filepath.Join(t.TempDir(), "non_existent.sock") + + client, err := New(nonExistentSocketPath) + assert.Error(t, err) + assert.Nil(t, client) + }) +} diff --git a/cmd/oneauth/rpcserver/listener.go b/cmd/oneauth/rpcserver/listener.go index acf937a..49353cc 100644 --- a/cmd/oneauth/rpcserver/listener.go +++ b/cmd/oneauth/rpcserver/listener.go @@ -2,21 +2,21 @@ package rpcserver import ( "context" - "fmt" + "errors" "net" - "net/http" + "net/rpc/jsonrpc" "os" - "time" + "strings" ) -func (s *RPCServer) ListenAndServe(_ context.Context, socketPath string) error { +func (s *RPCServer) ListenAndServe(ctx context.Context, socketPath string) error { defer func() { if _, err := os.Stat(socketPath); err == nil { os.Remove(socketPath) } }() - s.log.Println("listening rpc on", socketPath) + s.log.Println("listening json-rpc on", socketPath) listener, err := net.Listen("unix", socketPath) if err != nil { @@ -29,24 +29,28 @@ func (s *RPCServer) ListenAndServe(_ context.Context, socketPath string) error { return err } - mux := http.NewServeMux() - - mux.HandleFunc("/", func(w http.ResponseWriter, _ *http.Request) { - fmt.Fprintf(w, "hello world from oneauth agent") - }) - - server := &http.Server{ - Handler: mux, - ReadHeaderTimeout: 2 * time.Second, - } - s.mu.Lock() - s.server = server + s.listener = listener s.mu.Unlock() - if err := server.Serve(listener); err != nil && err != http.ErrServerClosed { - return err + go func() { + <-ctx.Done() + listener.Close() + }() + + for { + conn, err := listener.Accept() + if err != nil { + if isNetworkClosedError(err) || errors.Is(err, net.ErrClosed) { + return nil + } + return err + } + + go s.rpcServer.ServeCodec(jsonrpc.NewServerCodec(conn)) } +} - return nil +func isNetworkClosedError(err error) bool { + return err != nil && strings.Contains(err.Error(), "use of closed network connection") } diff --git a/cmd/oneauth/rpcserver/listener_test.go b/cmd/oneauth/rpcserver/listener_test.go index 3294343..682de3c 100644 --- a/cmd/oneauth/rpcserver/listener_test.go +++ b/cmd/oneauth/rpcserver/listener_test.go @@ -2,7 +2,6 @@ package rpcserver import ( "context" - "net/http" "os" "path/filepath" "testing" @@ -15,9 +14,7 @@ import ( func TestListenAndServe(t *testing.T) { t.Run("InvalidSocketPath", func(t *testing.T) { - rpcServer := &RPCServer{ - log: logrus.New(), - } + rpcServer := New(nil, logrus.New()) // Use invalid path (directory that doesn't exist) invalidPath := "/nonexistent/directory/socket" @@ -34,9 +31,7 @@ func TestListenAndServe(t *testing.T) { socketPath := filepath.Join(tempDir, "test.sock") - rpcServer := &RPCServer{ - log: logrus.New(), - } + rpcServer := New(nil, logrus.New()) // Run in goroutine since ListenAndServe blocks errChan := make(chan error) @@ -57,8 +52,10 @@ func TestListenAndServe(t *testing.T) { // Wait for completion select { case err := <-errChan: - // Should complete without error or with ErrServerClosed - assert.True(t, err == nil || err == http.ErrServerClosed) + // Should complete gracefully (nil) or with listener closed error + if err != nil { + assert.Contains(t, err.Error(), "use of closed network connection") + } case <-time.After(5 * time.Second): t.Fatal("ListenAndServe did not complete in time") } @@ -79,9 +76,7 @@ func TestListenAndServeCleanup(t *testing.T) { assert.NoError(t, err) file.Close() - rpcServer := &RPCServer{ - log: logrus.New(), - } + rpcServer := New(nil, logrus.New()) // Run in goroutine errChan := make(chan error) @@ -97,8 +92,9 @@ func TestListenAndServeCleanup(t *testing.T) { // Wait for completion select { - case <-errChan: - // Should complete + case err := <-errChan: + // Should complete with an error (expected since we pre-created a file) + assert.Error(t, err) case <-time.After(5 * time.Second): t.Fatal("ListenAndServe did not complete in time") } @@ -107,81 +103,53 @@ func TestListenAndServeCleanup(t *testing.T) { func TestListenAndServeHTTPHandling(t *testing.T) { t.Run("HTTPMuxSetup", func(t *testing.T) { - // Create temporary directory - tempDir, err := os.MkdirTemp("", "rpcserver_test") - assert.NoError(t, err) - defer os.RemoveAll(tempDir) - - socketPath := filepath.Join(tempDir, "test.sock") - - rpcServer := &RPCServer{ - log: logrus.New(), - } - - // Run in goroutine - errChan := make(chan error) - go func() { - errChan <- rpcServer.ListenAndServe(context.Background(), socketPath) - }() - - // Give it time to start - time.Sleep(100 * time.Millisecond) - - // Verify server was created - server := rpcServer.GetServer() - assert.NotNil(t, server) - assert.NotNil(t, server.Handler) - - // Shutdown the server - rpcServer.Shutdown() + testBasicServerStartup(t, "HTTPMuxSetup") + }) +} - // Wait for completion - select { - case <-errChan: - // Should complete - case <-time.After(5 * time.Second): - t.Fatal("ListenAndServe did not complete in time") - } +func TestListenAndServeJSONRPC(t *testing.T) { + t.Run("JSONRPCSetup", func(t *testing.T) { + testBasicServerStartup(t, "JSONRPCSetup") }) } -func TestListenAndServeTimeout(t *testing.T) { - t.Run("ReadHeaderTimeout", func(t *testing.T) { - // Create temporary directory - tempDir, err := os.MkdirTemp("", "rpcserver_test") - assert.NoError(t, err) - defer os.RemoveAll(tempDir) +// testBasicServerStartup is a helper function to test basic server startup and shutdown +func testBasicServerStartup(t *testing.T, testName string) { + // Create temporary directory + tempDir, err := os.MkdirTemp("", "rpcserver_test") + assert.NoError(t, err) + defer os.RemoveAll(tempDir) - socketPath := filepath.Join(tempDir, "test.sock") + socketPath := filepath.Join(tempDir, "test.sock") - rpcServer := &RPCServer{ - log: logrus.New(), - } + rpcServer := New(nil, logrus.New()) - // Run in goroutine - errChan := make(chan error) - go func() { - errChan <- rpcServer.ListenAndServe(context.Background(), socketPath) - }() + // Run in goroutine + errChan := make(chan error) + go func() { + errChan <- rpcServer.ListenAndServe(context.Background(), socketPath) + }() - // Give it time to start - time.Sleep(100 * time.Millisecond) + // Give it time to start + time.Sleep(100 * time.Millisecond) - // Verify timeout is set - server := rpcServer.GetServer() - assert.Equal(t, 2*time.Second, server.ReadHeaderTimeout) + // Verify RPC server is available + rpcSrv := rpcServer.GetRPCServer() + assert.NotNil(t, rpcSrv) - // Shutdown the server - rpcServer.Shutdown() + // Shutdown the server + rpcServer.Shutdown() - // Wait for completion - select { - case <-errChan: - // Should complete - case <-time.After(5 * time.Second): - t.Fatal("ListenAndServe did not complete in time") + // Wait for completion + select { + case err := <-errChan: + // Should complete gracefully (nil) or with listener closed error + if err != nil { + assert.Contains(t, err.Error(), "use of closed network connection") } - }) + case <-time.After(5 * time.Second): + t.Fatalf("%s: ListenAndServe did not complete in time", testName) + } } func TestListenAndServePermissions(t *testing.T) { @@ -193,9 +161,7 @@ func TestListenAndServePermissions(t *testing.T) { socketPath := filepath.Join(tempDir, "test.sock") - rpcServer := &RPCServer{ - log: logrus.New(), - } + rpcServer := New(nil, logrus.New()) // Run in goroutine errChan := make(chan error) @@ -216,8 +182,11 @@ func TestListenAndServePermissions(t *testing.T) { // Wait for completion select { - case <-errChan: - // Should complete + case err := <-errChan: + // Should complete gracefully (nil) or with listener closed error + if err != nil { + assert.Contains(t, err.Error(), "use of closed network connection") + } case <-time.After(5 * time.Second): t.Fatal("ListenAndServe did not complete in time") } @@ -233,9 +202,7 @@ func TestListenAndServeContext(t *testing.T) { socketPath := filepath.Join(tempDir, "test.sock") - rpcServer := &RPCServer{ - log: logrus.New(), - } + rpcServer := New(nil, logrus.New()) ctx := context.Background() @@ -253,8 +220,11 @@ func TestListenAndServeContext(t *testing.T) { // Wait for completion select { - case <-errChan: - // Should complete + case err := <-errChan: + // Should complete gracefully (nil) or with listener closed error + if err != nil { + assert.Contains(t, err.Error(), "use of closed network connection") + } case <-time.After(5 * time.Second): t.Fatal("ListenAndServe did not complete in time") } @@ -263,9 +233,7 @@ func TestListenAndServeContext(t *testing.T) { func TestListenAndServeErrorHandling(t *testing.T) { t.Run("ListenerError", func(t *testing.T) { - rpcServer := &RPCServer{ - log: logrus.New(), - } + rpcServer := New(nil, logrus.New()) // Use invalid path err := rpcServer.ListenAndServe(context.Background(), "/invalid/path/socket") @@ -273,9 +241,7 @@ func TestListenAndServeErrorHandling(t *testing.T) { }) t.Run("ChmodError", func(t *testing.T) { - rpcServer := &RPCServer{ - log: logrus.New(), - } + rpcServer := New(nil, logrus.New()) // Use a clearly invalid socket path err := rpcServer.ListenAndServe(context.Background(), "/proc/invalid/socket") @@ -307,8 +273,8 @@ func TestListenAndServeIntegration(t *testing.T) { time.Sleep(100 * time.Millisecond) // Verify everything is set up - server := rpcServer.GetServer() - assert.NotNil(t, server) + rpcSrv := rpcServer.GetRPCServer() + assert.NotNil(t, rpcSrv) assert.NotNil(t, rpcServer.SSHAgent) assert.NotNil(t, rpcServer.log) @@ -317,8 +283,11 @@ func TestListenAndServeIntegration(t *testing.T) { // Wait for completion select { - case <-errChan: - // Should complete + case err := <-errChan: + // Should complete gracefully (nil) or with listener closed error + if err != nil { + assert.Contains(t, err.Error(), "use of closed network connection") + } case <-time.After(5 * time.Second): t.Fatal("ListenAndServe did not complete in time") } diff --git a/cmd/oneauth/rpcserver/server.go b/cmd/oneauth/rpcserver/server.go index 5b7333c..4968ca8 100644 --- a/cmd/oneauth/rpcserver/server.go +++ b/cmd/oneauth/rpcserver/server.go @@ -1,8 +1,7 @@ package rpcserver import ( - "context" - "net/http" + "net/rpc" "sync" "time" @@ -11,35 +10,40 @@ import ( ) type RPCServer struct { - SSHAgent *sshagent.SSHAgent - server *http.Server - log *logrus.Logger - mu sync.RWMutex + SSHAgent *sshagent.SSHAgent + rpcServer *rpc.Server + log *logrus.Logger + mu sync.RWMutex + listener interface{ Close() error } + startTime time.Time } func New(sshAgent *sshagent.SSHAgent, log *logrus.Logger) *RPCServer { - return &RPCServer{ - SSHAgent: sshAgent, - log: log, + rpcServer := rpc.NewServer() + s := &RPCServer{ + SSHAgent: sshAgent, + rpcServer: rpcServer, + log: log, + startTime: time.Now(), } + + rpcServer.Register(&AgentService{server: s}) + + return s } func (s *RPCServer) Shutdown() { - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - - s.mu.RLock() - server := s.server - s.mu.RUnlock() + s.mu.Lock() + defer s.mu.Unlock() - if server != nil { - server.Shutdown(ctx) + if s.listener != nil { + s.listener.Close() + s.listener = nil } } -// GetServer returns the HTTP server instance (for testing) -func (s *RPCServer) GetServer() *http.Server { +func (s *RPCServer) GetRPCServer() *rpc.Server { s.mu.RLock() defer s.mu.RUnlock() - return s.server + return s.rpcServer } diff --git a/cmd/oneauth/rpcserver/server_test.go b/cmd/oneauth/rpcserver/server_test.go index 3664142..b3dada5 100644 --- a/cmd/oneauth/rpcserver/server_test.go +++ b/cmd/oneauth/rpcserver/server_test.go @@ -1,9 +1,7 @@ package rpcserver import ( - "net/http" "testing" - "time" "github.com/sirupsen/logrus" "github.com/stretchr/testify/assert" @@ -23,7 +21,7 @@ func TestNew(t *testing.T) { assert.NotNil(t, rpcServer) assert.Equal(t, sshAgent, rpcServer.SSHAgent) assert.Equal(t, log, rpcServer.log) - assert.Nil(t, rpcServer.GetServer()) + assert.NotNil(t, rpcServer.GetRPCServer()) }) } @@ -37,7 +35,7 @@ func TestRPCServerType(t *testing.T) { // Verify fields exist assert.NotNil(t, &rpcServer.SSHAgent) - assert.NotNil(t, rpcServer.GetServer) + assert.NotNil(t, rpcServer.GetRPCServer) assert.NotNil(t, &rpcServer.log) }) } @@ -53,8 +51,8 @@ func TestRPCServerFields(t *testing.T) { assert.Equal(t, sshAgent, rpcServer.SSHAgent) assert.Equal(t, log, rpcServer.log) - // Server should initially be nil - assert.Nil(t, rpcServer.GetServer()) + // RPC Server should be initialized + assert.NotNil(t, rpcServer.GetRPCServer()) }) } @@ -68,29 +66,37 @@ func TestShutdown(t *testing.T) { }) }) - t.Run("ShutdownWithServer", func(t *testing.T) { + t.Run("ShutdownWithListener", func(t *testing.T) { + mockListener := &mockListener{} rpcServer := &RPCServer{ - server: &http.Server{}, + listener: mockListener, } - // Should not panic when server exists + // Should not panic when listener exists assert.NotPanics(t, func() { rpcServer.Shutdown() }) + assert.True(t, mockListener.closed) }) } -func TestShutdownTimeout(t *testing.T) { - t.Run("TimeoutContextCreation", func(t *testing.T) { - rpcServer := &RPCServer{} +type mockListener struct { + closed bool +} + +func (m *mockListener) Close() error { + m.closed = true + return nil +} - // Test that shutdown creates proper context - start := time.Now() - rpcServer.Shutdown() - elapsed := time.Since(start) +func TestShutdownSpeed(t *testing.T) { + t.Run("QuickShutdown", func(t *testing.T) { + rpcServer := &RPCServer{} - // Should complete quickly when server is nil - assert.Less(t, elapsed, 100*time.Millisecond) + // Test that shutdown completes quickly + assert.NotPanics(t, func() { + rpcServer.Shutdown() + }) }) } @@ -104,7 +110,7 @@ func TestRPCServerConstruction(t *testing.T) { // Verify all parameters are set assert.NotNil(t, rpcServer.SSHAgent) assert.NotNil(t, rpcServer.log) - assert.Nil(t, rpcServer.GetServer()) + assert.NotNil(t, rpcServer.GetRPCServer()) }) t.Run("WithNilParameters", func(t *testing.T) { @@ -115,6 +121,7 @@ func TestRPCServerConstruction(t *testing.T) { assert.NotNil(t, rpcServer) assert.Nil(t, rpcServer.SSHAgent) assert.Nil(t, rpcServer.log) + assert.NotNil(t, rpcServer.GetRPCServer()) }) } @@ -133,26 +140,18 @@ func TestRPCServerMethods(t *testing.T) { } func TestRPCServerShutdownBehavior(t *testing.T) { - t.Run("ShutdownCancellation", func(t *testing.T) { - // Create a server with a mock HTTP server - mockServer := &http.Server{} + t.Run("ShutdownCompletion", func(t *testing.T) { + // Create a server with a mock listener + mockListener := &mockListener{} rpcServer := &RPCServer{ - server: mockServer, + listener: mockListener, } - // Test shutdown doesn't hang - done := make(chan bool) - go func() { + // Test shutdown completes + assert.NotPanics(t, func() { rpcServer.Shutdown() - done <- true - }() - - select { - case <-done: - // Success - shutdown completed - case <-time.After(10 * time.Second): - t.Fatal("Shutdown took too long") - } + }) + assert.True(t, mockListener.closed) }) } @@ -168,8 +167,9 @@ func TestRPCServerEdgeCases(t *testing.T) { }) t.Run("RepeatedShutdown", func(t *testing.T) { + mockListener := &mockListener{} rpcServer := &RPCServer{ - server: &http.Server{}, + listener: mockListener, } // Multiple shutdowns should not panic @@ -178,5 +178,6 @@ func TestRPCServerEdgeCases(t *testing.T) { rpcServer.Shutdown() rpcServer.Shutdown() }) + assert.True(t, mockListener.closed) }) } diff --git a/cmd/oneauth/rpcserver/service.go b/cmd/oneauth/rpcserver/service.go new file mode 100644 index 0000000..352eea0 --- /dev/null +++ b/cmd/oneauth/rpcserver/service.go @@ -0,0 +1,57 @@ +package rpcserver + +import ( + "fmt" + "os" + "time" + + "github.com/vitalvas/oneauth/internal/buildinfo" + "github.com/vitalvas/oneauth/internal/yubikey" +) + +// AgentService provides JSON-RPC methods for the SSH agent +type AgentService struct { + server *RPCServer +} + +// InfoArgs represents arguments for the Info method +type InfoArgs struct{} + +// InfoKey represents YubiKey information +type InfoKey struct { + Name string `json:"name"` + Serial string `json:"serial"` + Version string `json:"version"` +} + +// InfoReply represents the response from the Info method +type InfoReply struct { + Pid int `json:"pid"` + Keys []InfoKey `json:"keys"` + Version string `json:"version"` + Uptime string `json:"uptime"` +} + +// Info returns information about the running agent +func (s *AgentService) Info(_ *InfoArgs, reply *InfoReply) error { + reply.Pid = os.Getpid() + reply.Version = buildinfo.FormattedVersion() + + uptime := time.Since(s.server.startTime) + reply.Uptime = uptime.Truncate(time.Second).String() + + cards, err := yubikey.Cards() + if err != nil { + reply.Keys = []InfoKey{} + } else { + for _, card := range cards { + reply.Keys = append(reply.Keys, InfoKey{ + Name: card.String(), + Serial: fmt.Sprintf("%d", card.Serial), + Version: card.Version, + }) + } + } + + return nil +} diff --git a/cmd/oneauth/rpcserver/service_test.go b/cmd/oneauth/rpcserver/service_test.go new file mode 100644 index 0000000..c753c79 --- /dev/null +++ b/cmd/oneauth/rpcserver/service_test.go @@ -0,0 +1,47 @@ +package rpcserver + +import ( + "os" + "testing" + "time" + + "github.com/stretchr/testify/assert" +) + +func TestAgentServiceInfo(t *testing.T) { + t.Run("InfoMethod", func(t *testing.T) { + // Create a mock RPC server with start time + server := &RPCServer{ + startTime: time.Now().Add(-1 * time.Hour), // 1 hour ago + } + service := &AgentService{server: server} + args := &InfoArgs{} + reply := &InfoReply{} + + err := service.Info(args, reply) + assert.NoError(t, err) + assert.Equal(t, os.Getpid(), reply.Pid) + // Version might be empty in tests, but should not be nil + assert.NotNil(t, reply.Version) + assert.NotEmpty(t, reply.Uptime) + assert.Contains(t, reply.Uptime, "h") // Should contain hour indicator + }) + + t.Run("InfoReplyStructure", func(t *testing.T) { + reply := &InfoReply{ + Pid: 123, + Version: "1.0.0", + Uptime: "1h30m0s", + Keys: []InfoKey{}, + } + assert.Equal(t, 123, reply.Pid) + assert.Equal(t, "1.0.0", reply.Version) + assert.Equal(t, "1h30m0s", reply.Uptime) + assert.NotNil(t, reply.Keys) + }) + + t.Run("InfoArgsStructure", func(t *testing.T) { + args := &InfoArgs{} + assert.NotNil(t, args) + }) +} diff --git a/cmd/oneauth/service/service_darwin.go b/cmd/oneauth/service/service_darwin.go index 10e0abf..d940436 100644 --- a/cmd/oneauth/service/service_darwin.go +++ b/cmd/oneauth/service/service_darwin.go @@ -157,3 +157,16 @@ func writeServiceTemplate(exePath string, serviceFile *os.File) error { template.New("service").Parse(serviceTmpl), ).Execute(serviceFile, serviceInfo) } + +func IsRunning() bool { + output, err := callLaunchCtl("list", serviceName) + if err != nil { + return false + } + + if !strings.Contains(output, "PID") { + return false + } + + return true +} diff --git a/cmd/oneauth/service/service_linux.go b/cmd/oneauth/service/service_linux.go index 0d96ae3..c5d63b2 100644 --- a/cmd/oneauth/service/service_linux.go +++ b/cmd/oneauth/service/service_linux.go @@ -11,3 +11,7 @@ func Uninstal() error { func Restart() error { return ErrNotImplemented } + +func IsRunning() bool { + return false +} diff --git a/cmd/oneauth/sshagent/agent.go b/cmd/oneauth/sshagent/agent.go index 52c296f..996b258 100644 --- a/cmd/oneauth/sshagent/agent.go +++ b/cmd/oneauth/sshagent/agent.go @@ -33,9 +33,15 @@ type Actions struct { } func New(serial uint32, log *logrus.Logger, config *config.Config) (*SSHAgent, error) { - yk, err := yubikey.OpenBySerial(serial) - if err != nil { - return nil, err + var yk *yubikey.Yubikey + var err error + + // Only try to open YubiKey if serial is not 0 (disabled mode) + if serial != 0 { + yk, err = yubikey.OpenBySerial(serial) + if err != nil { + return nil, err + } } contextLogger := log.WithFields(logrus.Fields{ @@ -46,7 +52,7 @@ func New(serial uint32, log *logrus.Logger, config *config.Config) (*SSHAgent, e actions: Actions{ BeforeSignHook: config.Keyring.BeforeSignHook, }, - yk: yk, + yk: yk, // Will be nil when YubiKey is disabled log: contextLogger, softKeys: keystore.New(config.Keyring.KeepKeySeconds), diff --git a/cmd/oneauth/sshagent/ask_pin.go b/cmd/oneauth/sshagent/ask_pin.go index 76a7da2..d64226c 100644 --- a/cmd/oneauth/sshagent/ask_pin.go +++ b/cmd/oneauth/sshagent/ask_pin.go @@ -2,6 +2,7 @@ package sshagent import ( "errors" + "fmt" "github.com/vitalvas/oneauth/internal/keyring" ) @@ -11,6 +12,10 @@ var ( ) func (a *SSHAgent) askPINPrompt() (string, error) { + if a.yk == nil { + return "", fmt.Errorf("no yubikey available for PIN prompt") + } + pin, err := keyring.Get(keyring.GetYubikeyAccount(a.yk.Serial, "pin")) if err == nil { diff --git a/internal/buildinfo/version.go b/internal/buildinfo/version.go index 1278f79..a81f5f6 100644 --- a/internal/buildinfo/version.go +++ b/internal/buildinfo/version.go @@ -4,3 +4,12 @@ var ( Version string Commit string ) + +// FormattedVersion returns the version with commit hash if available +func FormattedVersion() string { + version := Version + if Commit != "" && len(Commit) >= 8 { + version += "-" + Commit[:8] + } + return version +} diff --git a/internal/buildinfo/version_test.go b/internal/buildinfo/version_test.go new file mode 100644 index 0000000..19eb487 --- /dev/null +++ b/internal/buildinfo/version_test.go @@ -0,0 +1,67 @@ +package buildinfo + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestFormattedVersion(t *testing.T) { + // Save original values + originalVersion := Version + originalCommit := Commit + + // Restore original values after test + defer func() { + Version = originalVersion + Commit = originalCommit + }() + + t.Run("VersionOnly", func(t *testing.T) { + Version = "1.0.0" + Commit = "" + + result := FormattedVersion() + assert.Equal(t, "1.0.0", result) + }) + + t.Run("VersionWithCommit", func(t *testing.T) { + Version = "1.0.0" + Commit = "abcdef1234567890" + + result := FormattedVersion() + assert.Equal(t, "1.0.0-abcdef12", result) + }) + + t.Run("VersionWithShortCommit", func(t *testing.T) { + Version = "1.0.0" + Commit = "abc123" + + result := FormattedVersion() + assert.Equal(t, "1.0.0", result) // Should not include short commit + }) + + t.Run("VersionWithExactly8CharCommit", func(t *testing.T) { + Version = "1.0.0" + Commit = "abcdef12" + + result := FormattedVersion() + assert.Equal(t, "1.0.0-abcdef12", result) + }) + + t.Run("EmptyVersion", func(t *testing.T) { + Version = "" + Commit = "abcdef1234567890" + + result := FormattedVersion() + assert.Equal(t, "-abcdef12", result) + }) + + t.Run("BothEmpty", func(t *testing.T) { + Version = "" + Commit = "" + + result := FormattedVersion() + assert.Equal(t, "", result) + }) +} \ No newline at end of file diff --git a/internal/mock/ssh.go b/internal/mock/ssh.go index 8525f13..c987a8b 100644 --- a/internal/mock/ssh.go +++ b/internal/mock/ssh.go @@ -123,7 +123,7 @@ func (m *ChannelConn) Read(b []byte) (n int, err error) { if m.readPos >= len(m.readData) { return 0, io.EOF } - + n = copy(b, m.readData[m.readPos:]) m.readPos += n return n, nil @@ -184,4 +184,4 @@ func MakeRequestSlice(requests []*ssh.Request) <-chan *ssh.Request { } close(ch) return ch -} \ No newline at end of file +} diff --git a/internal/mock/ssh_test.go b/internal/mock/ssh_test.go index 1f29753..2abc706 100644 --- a/internal/mock/ssh_test.go +++ b/internal/mock/ssh_test.go @@ -20,7 +20,7 @@ func TestNewChannel(t *testing.T) { conn := NewChannelConn() requests := MakeRequestSlice([]*ssh.Request{}) data := []byte("test data") - + channel := NewSSHChannel("session"). WithConn(conn). WithRequests(requests). @@ -33,7 +33,7 @@ func TestNewChannel(t *testing.T) { t.Run("Accept", func(t *testing.T) { conn := NewChannelConn() requests := MakeRequestSlice([]*ssh.Request{}) - + channel := NewSSHChannel("session"). WithConn(conn). WithRequests(requests) @@ -57,7 +57,7 @@ func TestNewChannel(t *testing.T) { t.Run("Reject", func(t *testing.T) { channel := NewSSHChannel("session") - + err := channel.Reject(ssh.UnknownChannelType, "test rejection") assert.NoError(t, err) assert.True(t, channel.IsRejected()) @@ -77,7 +77,7 @@ func TestChannelConn(t *testing.T) { t.Run("Write", func(t *testing.T) { conn := NewChannelConn() data := []byte("test data") - + n, err := conn.Write(data) assert.NoError(t, err) assert.Equal(t, len(data), n) @@ -86,7 +86,7 @@ func TestChannelConn(t *testing.T) { t.Run("Read", func(t *testing.T) { conn := NewChannelConn().WithReadData([]byte("test data")) - + buf := make([]byte, 5) n, err := conn.Read(buf) assert.NoError(t, err) @@ -96,11 +96,11 @@ func TestChannelConn(t *testing.T) { t.Run("Close", func(t *testing.T) { conn := NewChannelConn() - + err := conn.Close() assert.NoError(t, err) assert.True(t, conn.IsClosed()) - + // Writing to closed connection should fail _, err = conn.Write([]byte("test")) assert.Error(t, err) @@ -108,7 +108,7 @@ func TestChannelConn(t *testing.T) { t.Run("SendRequest", func(t *testing.T) { conn := NewChannelConn() - + ok, err := conn.SendRequest("test", true, []byte("payload")) assert.NoError(t, err) assert.True(t, ok) @@ -121,16 +121,16 @@ func TestHelperFunctions(t *testing.T) { NewSSHChannel("session"), NewSSHChannel("direct-tcpip"), } - + ch := MakeNewChannelSlice(channels) - + // Read from channel newChan1 := <-ch assert.Equal(t, "session", newChan1.ChannelType()) - + newChan2 := <-ch assert.Equal(t, "direct-tcpip", newChan2.ChannelType()) - + // Channel should be closed _, ok := <-ch assert.False(t, ok) @@ -141,20 +141,20 @@ func TestHelperFunctions(t *testing.T) { {Type: "shell", WantReply: true}, {Type: "exec", WantReply: false}, } - + ch := MakeRequestSlice(requests) - + // Read from channel req1 := <-ch assert.Equal(t, "shell", req1.Type) assert.True(t, req1.WantReply) - + req2 := <-ch assert.Equal(t, "exec", req2.Type) assert.False(t, req2.WantReply) - + // Channel should be closed _, ok := <-ch assert.False(t, ok) }) -} \ No newline at end of file +} diff --git a/internal/tools/dir.go b/internal/tools/dir.go index abfcbe0..68d95d9 100644 --- a/internal/tools/dir.go +++ b/internal/tools/dir.go @@ -4,6 +4,7 @@ import "os" func MkDir(dir string, perm os.FileMode) error { stat, err := os.Stat(dir) + if os.IsNotExist(err) { if err = os.MkdirAll(dir, perm); err != nil { return err