You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在Flask应用中获取SageMaker调用的CustomAttributes值?

如何在SageMaker端点的Flask应用中获取CustomAttributes值?

我有一段调用跨账号SageMaker端点的AWS Lambda函数代码:

import os
import boto3
from CustomModules.Logger import setlogging

global logger
logger = setlogging()


def lambda_handler(event, context):
    '''
    We use this lambda to call the sagemaker endpoint
    which lives on a different account.
    To do so, we need to assume a different role with boto3.
    '''
    # grab environment variables
    ENDPOINT_NAME = os.environ['ENDPOINT']
    someaccount = 'testaccount'
    runtime = boto3.client('runtime.sagemaker')
    sts_connection = boto3.client('sts')
    acct_b = sts_connection.assume_role(
        RoleArn=someaccount,
        RoleSessionName="somerole"
    )
    ACCESS_KEY = mydict['Credentials']['AccessKeyId']
    SECRET_KEY = mydict['Credentials']['SecretAccessKey']
    SESSION_TOKEN = mydict['Credentials']['SessionToken']

    # once we have all the info, we open the client connection with the new creds
    runtime = boto3.client(
        'runtime.sagemaker',
        aws_access_key_id=ACCESS_KEY,
        aws_secret_access_key=SECRET_KEY,
        aws_session_token=SESSION_TOKEN,
    )
    input_data = event['body']
    client_name = 'myclient'
    res = runtime.invoke_endpoint(EndpointName=ENDPOINT_NAME,
                                  ContentType='application/json',
                                  Body=input_data,
                                  CustomAttributes=client_name,
                                  Accept='Accept'
                                  )
    response = {
        "statusCode": res['ResponseMetadata']['HTTPStatusCode'],
        "headers": res['ResponseMetadata']['HTTPHeaders'],
        "body": res['Body'].read().decode('utf-8')}
    return response


if __name__ == '__main__':
    lambda_handler({''}, {''})

我希望在该SageMaker端点中获取设置的CustomAttributes值,该端点使用Flask处理请求,代码如下:

import flask

app = flask.Flask(__name__)

@app.route('/invocations', methods=['POST'])
def transformation():
    # Get input JSON data and convert it to a DF
    input_json = flask.request.get_json()
    ### How can I get the CustomAttributes value here?

请问如何在上述Flask应用中获取CustomAttributes的值?我遗漏了什么步骤?


解决方案

1. 先修复Lambda代码中的变量错误

你的Lambda代码存在一个明显的错误:sts_connection.assume_role()的返回值存在acct_b变量中,但后续获取凭证时用了未定义的mydict,这会触发NameError。需要把相关代码改成:

ACCESS_KEY = acct_b['Credentials']['AccessKeyId']
SECRET_KEY = acct_b['Credentials']['SecretAccessKey']
SESSION_TOKEN = acct_b['Credentials']['SessionToken']

2. 在Flask应用中读取CustomAttributes

AWS SageMaker会将你调用invoke_endpoint时传入的CustomAttributes值,通过HTTP请求头X-Amzn-SageMaker-Custom-Attributes传递给端点的Flask服务。你只需要在Flask的请求处理函数中读取这个请求头即可:

import flask

app = flask.Flask(__name__)

@app.route('/invocations', methods=['POST'])
def transformation():
    # Get input JSON data and convert it to a DF
    input_json = flask.request.get_json()
    
    # 获取CustomAttributes值
    custom_attributes = flask.request.headers.get('X-Amzn-SageMaker-Custom-Attributes')
    # 现在custom_attributes就是你在Lambda中设置的'client_name'值,即'myclient'
    
    # 后续业务逻辑...

额外说明:处理多键值对的CustomAttributes

如果你的CustomAttributes是类似client=myclient,env=prod这样的多键值对格式,可以自行解析成字典:

def parse_custom_attrs(attr_str):
    attr_dict = {}
    if attr_str:
        for item in attr_str.split(','):
            key, value = item.split('=', 1)
            attr_dict[key.strip()] = value.strip()
    return attr_dict

# 在请求处理函数中调用
custom_attrs_dict = parse_custom_attrs(flask.request.headers.get('X-Amzn-SageMaker-Custom-Attributes'))
# 例如custom_attrs_dict.get('client')就能拿到'myclient'

内容的提问来源于stack exchange,提问作者hadron

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.16 13:13:13