This is preparatory patch to use tsm measurement registers in IMA.
Introduce tsm_default_tm() and tsm_mr_read()/write() APIs
to read and extend the tsm measurement registers in IMA.

Since IMA is supported only when it's built as built-in,
export those symbols only tsm-mr is built as built-in.

Signed-off-by: Yeoreum Yun <[email protected]>
---
 drivers/virt/coco/guest/tsm-mr.c | 159 +++++++++++++++++++++++++++++++++------
 include/linux/tsm-mr.h           |  26 +++++++
 2 files changed, 161 insertions(+), 24 deletions(-)

diff --git a/drivers/virt/coco/guest/tsm-mr.c b/drivers/virt/coco/guest/tsm-mr.c
index 657b9c5739d0..9e721348be8d 100644
--- a/drivers/virt/coco/guest/tsm-mr.c
+++ b/drivers/virt/coco/guest/tsm-mr.c
@@ -7,9 +7,15 @@
 #include <linux/slab.h>
 #include <linux/sysfs.h>
 
+
 #define CREATE_TRACE_POINTS
 #include <trace/events/tsm_mr.h>
 
+#define TM_NUM_CTX     (64 * HASH_ALGO__LAST)
+
+DEFINE_IDR(tm_ctx_idr);
+static DEFINE_MUTEX(idr_lock);
+
 /*
  * struct tm_context - contains everything necessary to implement sysfs
  * attributes for MRs.
@@ -42,21 +48,16 @@ struct tm_context {
        struct bin_attribute mrs[];
 };
 
-static ssize_t tm_digest_read(struct file *filp, struct kobject *kobj,
-                             const struct bin_attribute *attr, char *buffer,
-                             loff_t off, size_t count)
+static ssize_t __tsm_mr_read(struct tm_context *ctx,
+                        const struct tsm_measurement_register *mr,
+                        char *buffer, loff_t off, size_t count)
 {
-       struct tm_context *ctx;
-       const struct tsm_measurement_register *mr;
        int rc;
 
-       ctx = attr->private;
        rc = down_read_interruptible(&ctx->rwsem);
        if (rc)
                return rc;
 
-       mr = &ctx->tm->mrs[attr - ctx->mrs];
-
        /*
         * @ctx->in_sync indicates if the MR cache is stale. It is a global
         * instead of a per-MR flag for simplicity, as most (if not all) archs
@@ -88,20 +89,11 @@ static ssize_t tm_digest_read(struct file *filp, struct 
kobject *kobj,
        return rc ?: count;
 }
 
-static ssize_t tm_digest_write(struct file *filp, struct kobject *kobj,
-                              const struct bin_attribute *attr, char *buffer,
-                              loff_t off, size_t count)
+static ssize_t __tsm_mr_write(struct tm_context *ctx,
+                        const struct tsm_measurement_register *mr,
+                        char *buffer, size_t count)
 {
-       struct tm_context *ctx;
-       const struct tsm_measurement_register *mr;
-       ssize_t rc;
-
-       /* partial writes are not supported */
-       if (off != 0 || count != attr->size)
-               return -EINVAL;
-
-       ctx = attr->private;
-       mr = &ctx->tm->mrs[attr - ctx->mrs];
+       int rc;
 
        rc = down_write_killable(&ctx->rwsem);
        if (rc)
@@ -119,6 +111,36 @@ static ssize_t tm_digest_write(struct file *filp, struct 
kobject *kobj,
        return rc ?: count;
 }
 
+static ssize_t tm_digest_read(struct file *filp, struct kobject *kobj,
+                             const struct bin_attribute *attr, char *buffer,
+                             loff_t off, size_t count)
+{
+       struct tm_context *ctx;
+       const struct tsm_measurement_register *mr;
+
+       ctx = attr->private;
+       mr = &ctx->tm->mrs[attr - ctx->mrs];
+
+       return __tsm_mr_read(ctx, mr, buffer, off, count);
+}
+
+static ssize_t tm_digest_write(struct file *filp, struct kobject *kobj,
+                              const struct bin_attribute *attr, char *buffer,
+                              loff_t off, size_t count)
+{
+       struct tm_context *ctx;
+       const struct tsm_measurement_register *mr;
+
+       /* partial writes are not supported */
+       if (off != 0 || count != attr->size)
+               return -EINVAL;
+
+       ctx = attr->private;
+       mr = &ctx->tm->mrs[attr - ctx->mrs];
+
+       return __tsm_mr_write(ctx, mr, buffer, count);
+}
+
 /**
  * tsm_mr_create_attribute_group() - creates an attribute group for measurement
  * registers (MRs)
@@ -138,8 +160,7 @@ static ssize_t tm_digest_write(struct file *filp, struct 
kobject *kobj,
  * * %-ENOMEM - Out of memory.
  */
 const struct attribute_group *
-tsm_mr_create_attribute_group(const struct tsm_measurements *tm)
-{
+tsm_mr_create_attribute_group(const struct tsm_measurements *tm) {
        size_t nlen;
 
        if (!tm || !tm->mrs)
@@ -230,6 +251,15 @@ tsm_mr_create_attribute_group(const struct 
tsm_measurements *tm)
        ctx->agrp.name = "measurements";
        ctx->agrp.bin_attrs = no_free_ptr(attrs);
        ctx->tm = tm;
+
+       guard(mutex)(&idr_lock);
+       ((struct tsm_measurements *)tm)->ctx_id = idr_alloc(&tm_ctx_idr, ctx, 0,
+                                                           TM_NUM_CTX, 
GFP_KERNEL);
+       if (tm->ctx_id < 0) {
+               kfree(ctx->agrp.bin_attrs);
+               return ERR_PTR(tm->ctx_id);
+       }
+
        return &no_free_ptr(ctx)->agrp;
 }
 EXPORT_SYMBOL_GPL(tsm_mr_create_attribute_group);
@@ -243,9 +273,90 @@ EXPORT_SYMBOL_GPL(tsm_mr_create_attribute_group);
  */
 void tsm_mr_free_attribute_group(const struct attribute_group *attr_grp)
 {
+       struct tm_context *ctx;
+
        if (!IS_ERR_OR_NULL(attr_grp)) {
+               ctx = container_of(attr_grp, struct tm_context, agrp);
+               scoped_guard(mutex, &idr_lock)
+                       idr_remove(&tm_ctx_idr, ctx->tm->ctx_id);
                kfree(attr_grp->bin_attrs);
-               kfree(container_of(attr_grp, struct tm_context, agrp));
+               kfree(ctx);
        }
 }
 EXPORT_SYMBOL_GPL(tsm_mr_free_attribute_group);
+
+#if defined(CONFIG_TSM_MEASUREMENTS)
+const struct tsm_measurements *tsm_default_tm(void)
+{
+       struct tm_context *ctx;
+       int next_id = 0;
+
+       guard(mutex)(&idr_lock);
+
+       ctx = idr_get_next(&tm_ctx_idr, &next_id);
+       if (!ctx)
+               return NULL;
+
+       return ctx->tm;
+}
+EXPORT_SYMBOL_GPL(tsm_default_tm);
+
+int tsm_mr_read(const struct tsm_measurements *tm, int idx,
+               u8 *digest, u32 digest_size)
+{
+       struct tm_context *ctx;
+       const struct tsm_measurement_register *mr;
+       int rc;
+
+       scoped_guard(mutex, &idr_lock)
+               ctx = idr_find(&tm_ctx_idr, tm->ctx_id);
+
+       if (IS_ERR_OR_NULL(ctx))
+               return -ENODEV;
+
+       if (!digest || (idx >= ctx->tm->nr_mrs) ||
+           (ctx->tm->mrs[idx].mr_size > digest_size) ||
+           !(ctx->tm->mrs[idx].mr_flags & TSM_MR_F_READABLE))
+               return -EINVAL;
+
+       mr = &ctx->tm->mrs[idx];
+
+       rc = __tsm_mr_read(ctx, mr, (char *)digest, 0, mr->mr_size);
+       if (rc < 0)
+               return rc;
+
+       return 0;
+}
+EXPORT_SYMBOL_GPL(tsm_mr_read);
+
+int tsm_mr_write(const struct tsm_measurements *tm, int idx,
+                u8 *digest, u32 digest_size)
+{
+       struct tm_context *ctx;
+       const struct tsm_measurement_register *mr;
+       int rc;
+
+       scoped_guard(mutex, &idr_lock)
+               ctx = idr_find(&tm_ctx_idr, tm->ctx_id);
+
+       if (IS_ERR_OR_NULL(ctx))
+               return -ENODEV;
+
+       if (!digest || (idx >= ctx->tm->nr_mrs) ||
+           !(ctx->tm->mrs[idx].mr_flags & TSM_MR_F_WRITABLE))
+               return -EINVAL;
+
+       /* partial writes are not supported */
+       if (ctx->tm->mrs[idx].mr_size != digest_size)
+               return -EINVAL;
+
+       mr = &ctx->tm->mrs[idx];
+
+       rc = __tsm_mr_write(ctx, mr, (char *)digest, mr->mr_size);
+       if (rc < 0)
+               return rc;
+
+       return 0;
+}
+EXPORT_SYMBOL_GPL(tsm_mr_write);
+#endif
diff --git a/include/linux/tsm-mr.h b/include/linux/tsm-mr.h
index 50a521f4ac97..43a0f761cd96 100644
--- a/include/linux/tsm-mr.h
+++ b/include/linux/tsm-mr.h
@@ -80,10 +80,36 @@ struct tsm_measurements {
        int (*refresh)(const struct tsm_measurements *tm);
        int (*write)(const struct tsm_measurements *tm,
                     const struct tsm_measurement_register *mr, const u8 *data);
+       int ctx_id;
 };
 
 const struct attribute_group *
 tsm_mr_create_attribute_group(const struct tsm_measurements *tm);
 void tsm_mr_free_attribute_group(const struct attribute_group *attr_grp);
 
+#if defined(CONFIG_TSM_MEASUREMENTS)
+const struct tsm_measurements *tsm_default_tm(void);
+int tsm_mr_read(const struct tsm_measurements *tm, int idx,
+               u8 *digest, u32 digest_size);
+int tsm_mr_write(const struct tsm_measurements *tm, int idx,
+                u8 *digest, u32 digest_size);
+#else
+static inline const struct tsm_measurements *tsm_default_tm(void)
+{
+       return NULL;
+}
+
+static inline int tsm_mr_read(const struct tsm_measurements *tm, int idx,
+                             u8 *digest, u32 digest_size)
+{
+       return 0;
+}
+
+static inline int tsm_mr_write(const struct tsm_measurements *tm, int idx,
+                              u8 *digest, u32 digest_size)
+{
+       return 0;
+}
+#endif
+
 #endif

-- 
2.43.0


Reply via email to