diff --git a/drivers/iommu/iommu.c b/drivers/iommu/iommu.c index d07fb960411..a6eae7e4127 100644 --- a/drivers/iommu/iommu.c +++ b/drivers/iommu/iommu.c @@ -32,7 +32,7 @@ #include static struct kset *iommu_group_kset; -static struct ida iommu_group_ida; +static struct idr iommu_group_idr; static struct mutex iommu_group_mutex; struct iommu_group { @@ -126,7 +126,7 @@ static void iommu_group_release(struct kobject *kobj) group->iommu_data_release(group->iommu_data); mutex_lock(&iommu_group_mutex); - ida_remove(&iommu_group_ida, group->id); + idr_remove(&iommu_group_idr, group->id); mutex_unlock(&iommu_group_mutex); kfree(group->name); @@ -167,22 +167,27 @@ struct iommu_group *iommu_group_alloc(void) mutex_lock(&iommu_group_mutex); again: - if (unlikely(0 == ida_pre_get(&iommu_group_ida, GFP_KERNEL))) { + if (unlikely(0 == idr_pre_get(&iommu_group_idr, GFP_KERNEL))) { kfree(group); mutex_unlock(&iommu_group_mutex); return ERR_PTR(-ENOMEM); } - if (-EAGAIN == ida_get_new(&iommu_group_ida, &group->id)) + ret = idr_get_new_above(&iommu_group_idr, group, 1, &group->id); + if (ret == -EAGAIN) goto again; - mutex_unlock(&iommu_group_mutex); + if (ret == -ENOSPC) { + kfree(group); + return ERR_PTR(ret); + } + ret = kobject_init_and_add(&group->kobj, &iommu_group_ktype, NULL, "%d", group->id); if (ret) { mutex_lock(&iommu_group_mutex); - ida_remove(&iommu_group_ida, group->id); + idr_remove(&iommu_group_idr, group->id); mutex_unlock(&iommu_group_mutex); kfree(group); return ERR_PTR(ret); @@ -425,6 +430,37 @@ struct iommu_group *iommu_group_get(struct device *dev) } EXPORT_SYMBOL_GPL(iommu_group_get); +/** + * iommu_group_find - Find and return the group based on the group name. + * Also increment the reference count. + * @name: the name of the group + * + * This function is called by iommu drivers and clients to get the group + * by the specified name. If found, the group is returned and the group + * reference is incremented, else NULL. + */ +struct iommu_group *iommu_group_find(const char *name) +{ + struct iommu_group *group; + int next = 0; + + mutex_lock(&iommu_group_mutex); + while ((group = idr_get_next(&iommu_group_idr, &next))) { + if (group->name) { + if (strcmp(group->name, name) == 0) + break; + } + ++next; + } + mutex_unlock(&iommu_group_mutex); + + if (group) + kobject_get(group->devices_kobj); + + return group; +} +EXPORT_SYMBOL_GPL(iommu_group_find); + /** * iommu_group_put - Decrement group reference * @group: the group to use @@ -888,7 +924,7 @@ static int __init iommu_init(void) { iommu_group_kset = kset_create_and_add("iommu_groups", NULL, kernel_kobj); - ida_init(&iommu_group_ida); + idr_init(&iommu_group_idr); mutex_init(&iommu_group_mutex); BUG_ON(!iommu_group_kset); diff --git a/include/linux/iommu.h b/include/linux/iommu.h index d4fa3aa3704..97b8e0ba0d8 100644 --- a/include/linux/iommu.h +++ b/include/linux/iommu.h @@ -138,6 +138,7 @@ extern void iommu_group_remove_device(struct device *dev); extern int iommu_group_for_each_dev(struct iommu_group *group, void *data, int (*fn)(struct device *, void *)); extern struct iommu_group *iommu_group_get(struct device *dev); +extern struct iommu_group *iommu_group_find(const char *name); extern void iommu_group_put(struct iommu_group *group); extern int iommu_group_register_notifier(struct iommu_group *group, struct notifier_block *nb); @@ -317,6 +318,11 @@ static inline struct iommu_group *iommu_group_get(struct device *dev) return NULL; } +static inline struct iommu_group *iommu_group_find(const char *name) +{ + return NULL; +} + static inline void iommu_group_put(struct iommu_group *group) { }