diff --git a/service/usbip/client.go b/service/usbip/client.go index b63d2c00f..e277cc5c0 100644 --- a/service/usbip/client.go +++ b/service/usbip/client.go @@ -99,13 +99,13 @@ func (c *ClientService) Close() error { return nil } -func (c *ClientService) runBusIDLoop(ctx context.Context, busid, description string) { +func (c *ClientService) runBusIDLoop(ctx context.Context, busid, description string, expected option.USBIPDeviceMatch) { for { if ctx.Err() != nil { return } c.assignment.SetActive(busid, true) - session, err := c.attemptAttach(ctx, busid) + session, err := c.attemptAttach(ctx, busid, expected) if err != nil { c.assignment.SetActive(busid, false) c.logger.Error("attach ", description, " (", busid, "): ", err) @@ -146,7 +146,7 @@ func (c *ClientService) runBusIDLoop(ctx context.Context, busid, description str } } -func (c *ClientService) attemptAttach(ctx context.Context, busid string) (AttachedSession, error) { +func (c *ClientService) attemptAttach(ctx context.Context, busid string, expected option.USBIPDeviceMatch) (AttachedSession, error) { conn, err := c.dialer.DialContext(ctx, N.NetworkTCP, c.serverAddr) if err != nil { return nil, E.Cause(err, "dial ", c.serverAddr) @@ -181,6 +181,10 @@ func (c *ClientService) attemptAttach(ctx context.Context, busid string) (Attach if err != nil { return nil, E.Cause(err, "read OP_REP_IMPORT body") } + err = c.verifyImportedDevice(busid, info, expected) + if err != nil { + return nil, err + } session, err := c.host.Attach(ctx, info, conn) if err != nil { return nil, err @@ -188,3 +192,56 @@ func (c *ClientService) attemptAttach(ctx context.Context, busid string) (Attach releaseConn = false return session, nil } + +// verifyImportedDevice checks the OP_REP_IMPORT identity against the +// rule the worker is serving. Bus ids are positional: after a server +// restart or a replug, the same busid can carry an arbitrary other +// device, and the worker retries its last busid every few seconds — +// without this check it would import whatever now sits on that port. +func (c *ClientService) verifyImportedDevice(busid string, info DeviceInfoTruncated, expected option.USBIPDeviceMatch) error { + replyBusID := info.BusIDString() + if replyBusID != busid { + return E.New("server attached ", replyBusID, " instead of ", busid) + } + key := DeviceKey{ + BusID: busid, + VendorID: info.IDVendor, + ProductID: info.IDProduct, + Serial: info.SerialString(), + } + if key.Serial == "" { + // Standard usbipd replies carry no serial; fall back to the + // control snapshot (the state the assignment was made from). + if snapshotKey, found := c.remoteDeviceKey(busid); found && + snapshotKey.VendorID == key.VendorID && snapshotKey.ProductID == key.ProductID { + key.Serial = snapshotKey.Serial + } + } + adjusted := expected + if adjusted.Serial != "" && key.Serial == "" { + // No serial available anywhere to compare against; enforcing it + // would reject every import from serial-less servers. + adjusted.Serial = "" + } + if !matches(adjusted, key) { + return E.New("imported device vid=", fmt.Sprintf("0x%04x", key.VendorID), + " pid=", fmt.Sprintf("0x%04x", key.ProductID), + " serial=", key.Serial, " does not match ", describeMatch(expected)) + } + return nil +} + +func (c *ClientService) remoteDeviceKey(busid string) (DeviceKey, bool) { + c.remoteAccess.Lock() + defer c.remoteAccess.Unlock() + device, found := c.remoteDevices[busid] + if !found { + return DeviceKey{}, false + } + return DeviceKey{ + BusID: device.BusID, + VendorID: device.VendorID, + ProductID: device.ProductID, + Serial: device.Serial, + }, true +} diff --git a/service/usbip/client_shared.go b/service/usbip/client_shared.go index 272c08e67..7e6be4f4c 100644 --- a/service/usbip/client_shared.go +++ b/service/usbip/client_shared.go @@ -394,10 +394,10 @@ func (c *ClientService) runAssignedWorker(worker *clientAssignedWorker) { if worker.target.fixedBusID != "" { match = option.USBIPDeviceMatch{BusID: worker.target.fixedBusID} } - go func(busid, description string) { + go func(busid, description string, expected option.USBIPDeviceMatch) { defer close(done) - c.runBusIDLoop(runCtx, busid, description) - }(desired, describeMatch(match)) + c.runBusIDLoop(runCtx, busid, description, expected) + }(desired, describeMatch(match), match) } } } @@ -426,7 +426,7 @@ func (c *ClientService) startRemoteBusIDWorkerLocked(busid string) { worker := &clientRemoteWorker{cancel: cancel} c.allWorkers[busid] = worker go func() { - c.runBusIDLoop(runCtx, busid, busid) + c.runBusIDLoop(runCtx, busid, busid, option.USBIPDeviceMatch{BusID: busid}) cancel() c.workerAccess.Lock() if c.allWorkers[busid] == worker {