1#include <drivers/usb/usb.h>
2#include <log.h>
3#include <mem/alloc.h>
4#include <mem/page.h>
5#include <mem/pmm.h>
6#include <mem/vmm.h>
7#include <sch/sched.h>
8#include <string.h>
9#include <thread/io_wait.h>
10#include <thread/thread.h>
11
12#include "internal.h"
13
14#ifdef DEBUG_USB
15LOG_SITE_DECLARE_DEFAULT(usb);
16#else
17LOG_SITE_DECLARE_DEFAULT(usb, .flags = LOG_SITE_LEVEL(LOG_ERROR));
18#endif
19
20LOG_HANDLE_DECLARE_DEFAULT(usb);
21
22enum usb_error usb_transfer_sync(enum usb_error (*fn)(struct usb_request *),
23 struct usb_request *request,
24 struct io_wait_token *tok) {
25 struct thread *curr = thread_get_current();
26 request->complete = usb_wake_waiter;
27 request->context = curr;
28
29 enum irql irql = irql_raise(new_level: IRQL_DISPATCH_LEVEL);
30
31 struct io_wait_token iowt = IO_WAIT_TOKEN_EMPTY;
32
33 /* Begin it regardless */
34 if (tok) {
35 io_wait_begin(out: tok, io_object: request->dev);
36 } else {
37 io_wait_begin(out: &iowt, io_object: request->dev);
38 }
39
40 enum usb_error ret = fn(request);
41 if (ret != USB_OK) {
42 usb_warn("ret != USB_OK");
43
44 thread_wake_unlocked(t: curr, r: THREAD_WAKE_REASON_BLOCKING_MANUAL,
45 wake_src: request->dev);
46
47 /* Always end it */
48 if (tok) {
49 io_wait_end(t: tok, act: IO_WAIT_END_NO_OP);
50 } else {
51 io_wait_end(t: &iowt, act: IO_WAIT_END_NO_OP);
52 }
53
54 irql_lower(old_level: irql);
55 return ret;
56 }
57
58 irql_lower(old_level: irql);
59
60 thread_yield_until_wake_match();
61
62 /* Only end if the wait is local OR it was fatal */
63 if (!tok) {
64 io_wait_end(t: &iowt, act: IO_WAIT_END_NO_OP);
65 } else if (tok && request->status != USB_OK) {
66 io_wait_end(t: tok, act: IO_WAIT_END_NO_OP);
67 }
68
69 return request->status;
70}
71
72void usb_wake_waiter(struct usb_request *rq) {
73 thread_wake_from_io_block(t: rq->context, wake_src: rq->dev);
74}
75
76void usb_destroy(struct usb_request *rq) {
77 kfree(rq);
78}
79
80uint8_t usb_construct_rq_bitmap(uint8_t transfer, uint8_t type, uint8_t recip) {
81 uint8_t bitmap = 0;
82 bitmap |= transfer << USB_REQUEST_TRANSFER_SHIFT;
83 bitmap |= type << USB_REQUEST_TYPE_SHIFT;
84 bitmap |= recip;
85 return bitmap;
86}
87
88static uint8_t usb_get_desc_bitmap(void) {
89 return usb_construct_rq_bitmap(USB_REQUEST_TRANS_DTH,
90 USB_REQUEST_TYPE_STANDARD,
91 USB_REQUEST_RECIPIENT_DEVICE);
92}
93
94enum usb_error usb_get_string_descriptor(struct usb_device *dev,
95 uint8_t string_idx, char *out,
96 size_t max_len) {
97 enum usb_error err = USB_OK;
98 if (!string_idx)
99 return USB_ERR_INVALID_ARGUMENT;
100
101 struct usb_controller *ctrl = dev->host;
102 uint8_t *desc = kmalloc_aligned(PAGE_SIZE, PAGE_SIZE);
103
104 struct usb_setup_packet setup = {
105 .bitmap_request_type = usb_get_desc_bitmap(),
106 .request = USB_RQ_CODE_GET_DESCRIPTOR,
107 .value = (USB_DESC_TYPE_STRING << USB_DESC_TYPE_SHIFT) | string_idx,
108 .index = 0,
109 .length = 255,
110 };
111
112 struct usb_request req = {
113 .setup = &setup,
114 .buffer = desc,
115 .dev = dev,
116 };
117
118 struct io_wait_token tok = IO_WAIT_TOKEN_EMPTY;
119 if ((err = usb_transfer_sync(fn: ctrl->ops->submit_control_transfer, request: &req,
120 tok: &tok)) != USB_OK)
121 return err;
122
123 uint8_t bLength = desc[0];
124 if (bLength < 2) {
125 io_wait_end(t: &tok, act: IO_WAIT_END_YIELD);
126 return USB_ERR_INVALID_ARGUMENT;
127 }
128
129 size_t out_idx = 0;
130 for (size_t i = 2; i < bLength && out_idx < (max_len - 1); i += 2) {
131 out[out_idx++] = (desc[i + 1] == 0) ? desc[i] : '?';
132 }
133 out[out_idx] = '\0';
134
135 kfree_aligned(desc);
136 io_wait_end(t: &tok, act: IO_WAIT_END_YIELD);
137 return err;
138}
139
140enum usb_error usb_get_device_descriptor(struct usb_device *dev) {
141 uint8_t *desc = kmalloc_aligned(PAGE_SIZE, PAGE_SIZE);
142
143 struct usb_setup_packet setup = {
144 .bitmap_request_type = usb_get_desc_bitmap(),
145 .request = USB_RQ_CODE_GET_DESCRIPTOR,
146 .value = (USB_DESC_TYPE_DEVICE << USB_DESC_TYPE_SHIFT),
147 .index = 0,
148 .length = 18,
149 };
150
151 struct usb_controller *ctrl = dev->host;
152
153 struct usb_request request = {
154 .setup = &setup,
155 .buffer = desc,
156 .dev = dev,
157 };
158
159 enum usb_error err;
160 if ((err = usb_transfer_sync(fn: ctrl->ops->submit_control_transfer, request: &request,
161 NULL)) != USB_OK) {
162 return err;
163 }
164
165 struct usb_device_descriptor *ddesc = (void *) desc;
166
167 usb_get_string_descriptor(dev, string_idx: ddesc->manufacturer, out: dev->manufacturer,
168 max_len: sizeof(dev->manufacturer));
169 usb_get_string_descriptor(dev, string_idx: ddesc->product, out: dev->product,
170 max_len: sizeof(dev->product));
171
172 dev->descriptor = ddesc;
173 return USB_OK;
174}
175
176static void match_interfaces(struct usb_driver *driver,
177 struct usb_device *dev) {
178 for (uint8_t i = 0; i < dev->num_interfaces; i++) {
179 struct usb_interface_descriptor *in = dev->interfaces[i];
180 bool class, subclass, proto;
181 class = driver->class_code == 0xFF || driver->class_code == in->class;
182 subclass = driver->subclass == 0xFF || driver->subclass == in->subclass;
183 proto = driver->protocol == 0xFF || driver->protocol == in->protocol;
184
185 bool everything_matches = class && subclass && proto;
186 if (everything_matches) {
187 if (driver->bringup) {
188 driver->bringup(dev);
189 dev->driver = driver;
190 dev->free = driver->free;
191 dev->teardown = driver->teardown;
192 return;
193 }
194 }
195 }
196}
197
198void usb_try_bind_driver(struct usb_device *dev) {
199 struct usb_driver *start = __skernel_usb_drivers;
200 struct usb_driver *end = __ekernel_usb_drivers;
201
202 for (struct usb_driver *d = start; d < end; d++)
203 match_interfaces(driver: d, dev);
204}
205
206static void
207usb_register_dev_interface(struct usb_device *dev,
208 struct usb_interface_descriptor *interface) {
209 struct usb_interface_descriptor *new_int =
210 kmalloc(sizeof(struct usb_interface_descriptor));
211 memcpy(new_int, interface, sizeof(struct usb_interface_descriptor));
212
213 size_t size = (dev->num_interfaces + 1) * sizeof(void *);
214
215 dev->interfaces = krealloc(dev->interfaces, size);
216 dev->interfaces[dev->num_interfaces++] = new_int;
217}
218
219static void usb_register_dev_ep(struct usb_device *dev,
220 struct usb_endpoint_descriptor *endpoint) {
221 struct usb_endpoint_descriptor *new_ep =
222 kmalloc(sizeof(struct usb_endpoint_descriptor));
223 memcpy(new_ep, endpoint, sizeof(struct usb_endpoint_descriptor));
224
225 struct usb_endpoint *ep =
226 kmalloc(sizeof(struct usb_endpoint), ALLOC_FLAGS_ZERO);
227
228 ep->type = USB_ENDPOINT_ATTR_TRANS_TYPE(endpoint->attributes);
229 ep->number = USB_ENDPOINT_ADDR_EP_NUM(endpoint->address);
230 ep->address = endpoint->address;
231 ep->max_packet_size = endpoint->max_packet_size;
232 ep->interval = endpoint->interval;
233
234 ep->in = USB_ENDPOINT_ADDR_EP_DIRECTION(endpoint->address);
235
236 size_t size = (dev->num_endpoints + 1) * sizeof(void *);
237 dev->endpoints = krealloc(dev->endpoints, size);
238 dev->endpoints[dev->num_endpoints++] = ep;
239}
240
241static void setup_config_descriptor(struct usb_device *dev, uint8_t *ptr,
242 uint8_t *end) {
243 while (ptr < end) {
244 uint8_t len = ptr[0];
245 uint8_t dtype = ptr[1];
246
247 if (dtype == USB_DESC_TYPE_INTERFACE) {
248 struct usb_interface_descriptor *iface = (void *) ptr;
249 usb_register_dev_interface(dev, interface: iface);
250 } else if (dtype == USB_DESC_TYPE_ENDPOINT) {
251 struct usb_endpoint_descriptor *epd = (void *) ptr;
252 usb_register_dev_ep(dev, endpoint: epd);
253 }
254
255 ptr += len;
256 }
257}
258
259enum usb_error usb_parse_config_descriptor(struct usb_device *dev) {
260 uint8_t *desc = kmalloc_aligned(PAGE_SIZE, PAGE_SIZE);
261
262 struct usb_setup_packet setup = {
263 .bitmap_request_type = usb_get_desc_bitmap(),
264 .request = USB_RQ_CODE_GET_DESCRIPTOR,
265 .value = (USB_DESC_TYPE_CONFIG << USB_DESC_TYPE_SHIFT),
266 .index = 0,
267 .length = 18,
268 };
269
270 struct usb_controller *ctrl = dev->host;
271
272 struct usb_request request = {
273 .setup = &setup,
274 .buffer = desc,
275 .dev = dev,
276 };
277
278 enum usb_error err;
279 struct io_wait_token iowt = IO_WAIT_TOKEN_EMPTY;
280 if ((err = usb_transfer_sync(fn: ctrl->ops->submit_control_transfer, request: &request,
281 tok: &iowt)) != USB_OK) {
282 kfree_aligned(desc);
283 return err;
284 }
285
286 struct usb_config_descriptor *cdesc = (void *) desc;
287 memcpy(&dev->config, cdesc, sizeof(struct usb_config_descriptor));
288
289 uint16_t total_len = cdesc->total_length;
290 setup.length = total_len;
291
292 if ((err = usb_transfer_sync(fn: ctrl->ops->submit_control_transfer, request: &request,
293 NULL)) != USB_OK) {
294 kfree_aligned(desc);
295 io_wait_end(t: &iowt, act: IO_WAIT_END_YIELD);
296 return err;
297 }
298
299 if ((err = usb_get_string_descriptor(dev, string_idx: cdesc->configuration,
300 out: dev->config_str,
301 max_len: sizeof(dev->config_str))) != USB_OK) {
302 kfree_aligned(desc);
303 io_wait_end(t: &iowt, act: IO_WAIT_END_YIELD);
304 return err;
305 }
306
307 setup_config_descriptor(dev, ptr: desc, end: desc + total_len);
308 kfree_aligned(desc);
309 io_wait_end(t: &iowt, act: IO_WAIT_END_YIELD);
310 return USB_OK;
311}
312
313enum usb_error usb_set_configuration(struct usb_device *dev) {
314 struct usb_controller *ctrl = dev->host;
315
316 uint8_t bitmap = usb_construct_rq_bitmap(USB_REQUEST_TRANS_HTD,
317 USB_REQUEST_TYPE_STANDARD,
318 USB_REQUEST_RECIPIENT_DEVICE);
319
320 struct usb_setup_packet set_cfg = {
321 .bitmap_request_type = bitmap,
322 .request = USB_RQ_CODE_SET_CONFIG,
323 .value = dev->config.configuration_value,
324 .index = 0,
325 .length = 0,
326 };
327
328 struct usb_request request = {
329 .setup = &set_cfg,
330 .buffer = NULL,
331 .dev = dev,
332 };
333
334 enum usb_error err;
335 if ((err = usb_transfer_sync(fn: ctrl->ops->submit_control_transfer, request: &request,
336 NULL)) != USB_OK) {
337 return err;
338 }
339
340 return USB_OK;
341}
342
343struct usb_interface_descriptor *usb_find_interface(struct usb_device *dev,
344 uint8_t class,
345 uint8_t subclass,
346 uint8_t protocol) {
347 for (size_t i = 0; i < dev->num_interfaces; i++) {
348 struct usb_interface_descriptor *intf = dev->interfaces[i];
349 if (intf->class == class && intf->subclass == subclass &&
350 intf->protocol == protocol)
351 return intf;
352 }
353 return NULL;
354}
355
356void usb_print_device(struct usb_device *dev) {
357 usb_log(LOG_INFO, "Found device '%s' manufactured by '%s' of type '%s'",
358 dev->product, dev->manufacturer, dev->config_str);
359}
360
361enum usb_error usb_init_device(struct usb_device *dev) {
362 if (!usb_device_get(obj: dev))
363 return USB_ERR_NO_DEVICE;
364
365 enum usb_error err = USB_OK;
366 usb_trace("get_device_descriptor");
367 if ((err = usb_get_device_descriptor(dev)) != USB_OK) {
368 goto out;
369 }
370
371 usb_trace("parse_config");
372 if ((err = usb_parse_config_descriptor(dev)) != USB_OK) {
373 goto out;
374 }
375
376 usb_trace("set_config");
377 if ((err = usb_set_configuration(dev)) != USB_OK) {
378 goto out;
379 }
380
381 usb_trace("configure_endpoint");
382 if ((err = dev->host->ops->configure_endpoint(dev)) != USB_OK) {
383 goto out;
384 }
385
386 usb_print_device(dev);
387
388 usb_try_bind_driver(dev);
389
390out:
391 if (err != USB_OK) {
392 kassert(thread_get_current()->perceived_prio_class ==
393 THREAD_PRIO_CLASS_TIMESHARE);
394 usb_trace("reset_slot");
395 dev->host->ops->reset_slot(dev);
396 }
397
398 usb_device_put(dev);
399 usb_trace("ok");
400 return err;
401}
402
403void usb_teardown_device(struct usb_device *dev) {
404 struct usb_driver *driver = dev->driver;
405
406 atomic_store_explicit(&dev->status, USB_DEV_DISCONNECTED,
407 memory_order_release);
408
409 if (driver && dev->teardown)
410 dev->teardown(dev);
411
412 dev->driver = NULL;
413 dev->free = NULL;
414 dev->teardown = NULL;
415
416 usb_device_put(dev);
417}
418
419void usb_free_device(struct usb_device *dev) {
420 if (dev->free)
421 dev->free(dev);
422
423 for (size_t i = 0; i < dev->num_interfaces; i++) {
424 struct usb_interface_descriptor *infdr = dev->interfaces[i];
425 kfree(infdr);
426 }
427 kfree(dev->interfaces);
428
429 for (size_t i = 0; i < dev->num_endpoints; i++) {
430 struct usb_endpoint *uep = dev->endpoints[i];
431 kfree(uep);
432 }
433 kfree(dev->endpoints);
434 kfree_aligned(dev->descriptor);
435}
436