Skip to content

Commit 3d67dc6

Browse files
committed
fix: harden wrapper agent startup and proxy fallback
1 parent 309e833 commit 3d67dc6

5 files changed

Lines changed: 159 additions & 35 deletions

File tree

README.md

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,7 @@
22

33
[![Actions Status](https://github.com/epithet-ssh/epithet/workflows/build/badge.svg)](https://github.com/epithet-ssh/epithet/actions) [![Go Reportcard](https://goreportcard.com/badge/github.com/epithet-ssh/epithet)](https://goreportcard.com/report/github.com/epithet-ssh/epithet)
44

5-
Epithet is an SSH certificate authority that replaces static authorized_keys with short-lived certificates (2-10 minutes). It creates on-demand SSH agents for each outbound connection, enabling real-time policy enforcement without touching your target hosts.
5+
Epithet is an SSH agent and certificate authority that replaces static authorized_keys with short-lived certificates (2-10 minutes). It creates on-demand SSH agents for each outbound connection, enabling real-time policy enforcement without touching your target hosts.
66

77
## Quick start
88

cmd/epithet/agent.go

Lines changed: 31 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -113,6 +113,7 @@ func (s *AgentStartCLI) runWrapper(parent *AgentCLI, logger *slog.Logger, tlsCfg
113113
hello, err := agent.ProbeUpstream(upstreamSocket)
114114
if err != nil {
115115
logger.Warn("failed to probe upstream agent", "error", err)
116+
upstreamSocket = ""
116117
} else if hello != nil {
117118
depth = hello.ChainDepth + 1
118119
brokerOpts = append(brokerOpts, broker.WithUpstream(upstreamSocket))
@@ -145,7 +146,9 @@ func (s *AgentStartCLI) runWrapper(parent *AgentCLI, logger *slog.Logger, tlsCfg
145146
go func() {
146147
brokerErr <- b.Serve(ctx)
147148
}()
148-
<-b.Ready()
149+
if err := waitForBrokerStartup(b, brokerErr); err != nil {
150+
return err
151+
}
149152

150153
// Generate SSH config.
151154
sshConfigPath := filepath.Join(tempDir, "ssh-config.conf")
@@ -158,25 +161,22 @@ func (s *AgentStartCLI) runWrapper(parent *AgentCLI, logger *slog.Logger, tlsCfg
158161
}
159162
}
160163

161-
// Create proxy listener if we have an upstream socket to proxy.
162-
if upstreamSocket != "" {
163-
proxySock := filepath.Join(tempDir, "proxy.sock")
164-
setup := func(p *agent.ProxyAgent) {
165-
p.RegisterExtension(agent.ExtensionHello, agent.HelloHandler(depth))
166-
p.RegisterExtension(agent.ExtensionAuth, agent.AuthHandler(func() (string, error) {
167-
return b.Authenticate(nil)
168-
}))
169-
}
170-
proxyListener := agent.NewProxyListener(logger, proxySock, upstreamSocket, setup)
171-
go func() {
172-
if err := proxyListener.Serve(ctx); err != nil && err != context.Canceled {
173-
logger.Error("proxy listener error", "error", err)
174-
}
175-
}()
176-
<-proxyListener.Ready()
177-
upstreamSocket = proxySock
178-
logger.Info("proxy agent listening", "socket", proxySock)
164+
proxySock := filepath.Join(tempDir, "proxy.sock")
165+
setup := func(p *agent.ProxyAgent) {
166+
p.RegisterExtension(agent.ExtensionHello, agent.HelloHandler(depth))
167+
p.RegisterExtension(agent.ExtensionAuth, agent.AuthHandler(func() (string, error) {
168+
return b.Authenticate(nil)
169+
}))
179170
}
171+
proxyListener := agent.NewProxyListener(logger, proxySock, upstreamSocket, setup)
172+
go func() {
173+
if err := proxyListener.Serve(ctx); err != nil && err != context.Canceled {
174+
logger.Error("proxy listener error", "error", err)
175+
}
176+
}()
177+
<-proxyListener.Ready()
178+
upstreamSocket = proxySock
179+
logger.Info("proxy agent listening", "socket", proxySock)
180180

181181
// Build child environment with updated SSH_AUTH_SOCK.
182182
childEnv := os.Environ()
@@ -223,6 +223,18 @@ func (s *AgentStartCLI) runWrapper(parent *AgentCLI, logger *slog.Logger, tlsCfg
223223
return nil
224224
}
225225

226+
func waitForBrokerStartup(b *broker.Broker, brokerErr <-chan error) error {
227+
select {
228+
case <-b.Ready():
229+
return nil
230+
case err := <-brokerErr:
231+
if err == nil || err == context.Canceled {
232+
return fmt.Errorf("broker exited before becoming ready")
233+
}
234+
return fmt.Errorf("broker failed to start: %w", err)
235+
}
236+
}
237+
226238
// setupBrokerWithOptions creates the broker, temp directories, and resolves auth.
227239
// When wrapperMode is true, uses a per-process temp directory to avoid collisions
228240
// between multiple wrapped shells. When false (daemon mode), uses a deterministic

cmd/epithet/agent_test.go

Lines changed: 48 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,16 @@
11
package main
22

33
import (
4+
"context"
5+
"io"
6+
"log/slog"
7+
"os"
8+
"path/filepath"
9+
"strings"
410
"testing"
11+
12+
"github.com/epithet-ssh/epithet/pkg/broker"
13+
"github.com/epithet-ssh/epithet/pkg/caclient"
514
)
615

716
func TestReplaceEnv_ExistingKey(t *testing.T) {
@@ -49,3 +58,42 @@ func TestReplaceEnv_EmptyEnv(t *testing.T) {
4958
t.Errorf("expected [KEY=value], got %v", result)
5059
}
5160
}
61+
62+
func TestWaitForBrokerStartup_ReturnsServeError(t *testing.T) {
63+
t.Parallel()
64+
65+
tempDir := t.TempDir()
66+
agentDir := filepath.Join(tempDir, "agent")
67+
if err := os.MkdirAll(agentDir, 0700); err != nil {
68+
t.Fatalf("mkdir agent dir: %v", err)
69+
}
70+
71+
endpoints := []caclient.CAEndpoint{{URL: "http://127.0.0.1", Priority: caclient.DefaultPriority}}
72+
caClient, err := caclient.New(endpoints)
73+
if err != nil {
74+
t.Fatalf("create CA client: %v", err)
75+
}
76+
77+
logger := slog.New(slog.NewTextHandler(io.Discard, nil))
78+
socketPath := filepath.Join(tempDir, strings.Repeat("sock", 40))
79+
b, err := broker.New(*logger, socketPath, "true", caClient, agentDir)
80+
if err != nil {
81+
t.Fatalf("create broker: %v", err)
82+
}
83+
84+
ctx, cancel := context.WithCancel(context.Background())
85+
defer cancel()
86+
87+
brokerErr := make(chan error, 1)
88+
go func() {
89+
brokerErr <- b.Serve(ctx)
90+
}()
91+
92+
err = waitForBrokerStartup(b, brokerErr)
93+
if err == nil {
94+
t.Fatal("expected startup error, got nil")
95+
}
96+
if !strings.Contains(err.Error(), "broker failed to start") {
97+
t.Fatalf("expected startup failure wrapper, got %v", err)
98+
}
99+
}

pkg/agent/proxy_listener.go

Lines changed: 15 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -89,25 +89,26 @@ func (p *ProxyListener) acceptLoop() {
8989
func (p *ProxyListener) serveConn(conn net.Conn) {
9090
defer conn.Close()
9191

92-
// Dial a fresh connection to the upstream agent for this client.
93-
upstreamConn, err := net.Dial("unix", p.upstreamSocketPath)
94-
if err != nil {
95-
p.log.Warn("failed to connect to upstream agent", "error", err)
96-
return
97-
}
98-
defer upstreamConn.Close()
92+
var upstream agent.ExtendedAgent
93+
if p.upstreamSocketPath == "" {
94+
// Preserve standard ssh-agent behavior even without an upstream socket.
95+
upstream = agent.NewKeyring().(agent.ExtendedAgent)
96+
} else {
97+
// Dial a fresh connection to the upstream agent for this client.
98+
upstreamConn, err := net.Dial("unix", p.upstreamSocketPath)
99+
if err != nil {
100+
p.log.Warn("failed to connect to upstream agent", "error", err)
101+
return
102+
}
103+
defer upstreamConn.Close()
99104

100-
upstream := agent.NewClient(upstreamConn)
101-
extUpstream, ok := upstream.(agent.ExtendedAgent)
102-
if !ok {
103-
p.log.Warn("upstream agent does not support extensions")
104-
return
105+
upstream = agent.NewClient(upstreamConn)
105106
}
106107

107-
proxy := NewProxyAgent(extUpstream)
108+
proxy := NewProxyAgent(upstream)
108109
p.setup(proxy)
109110

110-
err = agent.ServeAgent(proxy, conn)
111+
err := agent.ServeAgent(proxy, conn)
111112
if err != nil && err != io.EOF {
112113
p.log.Debug("proxy agent connection ended", "error", err)
113114
}

pkg/agent/proxy_listener_test.go

Lines changed: 64 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -165,6 +165,70 @@ func TestProxyListenerMultipleConnections(t *testing.T) {
165165
}
166166
}
167167

168+
func TestProxyListenerWithoutUpstream(t *testing.T) {
169+
proxySock := tempSocketPath(t)
170+
logger := slog.New(slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{Level: slog.LevelError}))
171+
172+
var authCalled atomic.Bool
173+
setup := func(p *ProxyAgent) {
174+
p.RegisterExtension(ExtensionHello, HelloHandler(0))
175+
p.RegisterExtension(ExtensionAuth, AuthHandler(func() (string, error) {
176+
authCalled.Store(true)
177+
return "local-token", nil
178+
}))
179+
}
180+
proxy := NewProxyListener(logger, proxySock, "", setup)
181+
182+
ctx, cancel := context.WithCancel(context.Background())
183+
defer cancel()
184+
185+
go proxy.Serve(ctx)
186+
<-proxy.Ready()
187+
188+
conn, err := net.Dial("unix", proxySock)
189+
if err != nil {
190+
t.Fatal(err)
191+
}
192+
defer conn.Close()
193+
194+
client := agent.NewClient(conn).(agent.ExtendedAgent)
195+
196+
resp, err := client.Extension(ExtensionHello, nil)
197+
if err != nil {
198+
t.Fatalf("hello failed: %v", err)
199+
}
200+
var hello HelloResponse
201+
if err := json.Unmarshal(stripSuccessByte(t, resp), &hello); err != nil {
202+
t.Fatalf("unmarshal hello: %v", err)
203+
}
204+
if hello.ChainDepth != 0 {
205+
t.Errorf("expected chain depth 0, got %d", hello.ChainDepth)
206+
}
207+
208+
resp, err = client.Extension(ExtensionAuth, nil)
209+
if err != nil {
210+
t.Fatalf("auth failed: %v", err)
211+
}
212+
var authResp AuthResponse
213+
if err := json.Unmarshal(stripSuccessByte(t, resp), &authResp); err != nil {
214+
t.Fatalf("unmarshal auth: %v", err)
215+
}
216+
if authResp.Token != "local-token" {
217+
t.Errorf("expected local token, got %q", authResp.Token)
218+
}
219+
if !authCalled.Load() {
220+
t.Fatal("expected auth handler to be called")
221+
}
222+
223+
keys, err := client.List()
224+
if err != nil {
225+
t.Fatalf("List failed: %v", err)
226+
}
227+
if len(keys) != 0 {
228+
t.Errorf("expected empty local keyring, got %d keys", len(keys))
229+
}
230+
}
231+
168232
func TestMultiHopChain(t *testing.T) {
169233
// Simulate: laptop → shell1 → shell2
170234
// Each hop has its own proxy listener wrapping the previous.
@@ -279,4 +343,3 @@ func tempSocketPath(t *testing.T) string {
279343
t.Cleanup(func() { os.Remove(path) })
280344
return path
281345
}
282-

0 commit comments

Comments
 (0)