"""Connector for communicating with the Switchboard API.""" import time import boto3 import jwt from sgqlc.endpoint.http import HTTPEndpoint import config # noqa sts = boto3.client("sts") def get_switchboard_credentials(): """Get credentials for interacting with Switchboard AWS.""" assumed_role = sts.assume_role( RoleArn=config.SWITCHBOARD_ROLE_ARN, RoleSessionName="switchboard-graphql" ) credentials = assumed_role["Credentials"] return { "aws_access_key_id": credentials["AccessKeyId"], "aws_secret_access_key": credentials["SecretAccessKey"], "aws_session_token": credentials["SessionToken"], } def get_secret_value(arn, credentials): """Get Amazon Secrets Manager secret value. Args: arn (str): The ARN/name of the secrets manager secret credentials (dict): AWS credentials to use Returns: Queue: Instance of Amazon SQS queue. """ secretsmanager = boto3.client("secretsmanager", **credentials) secret = secretsmanager.get_secret_value( SecretId=arn, ) return secret GET_PRODUCT_BY_ID_QUERY = """ query($localId: String!, $system: System!) { getProductByLocalId(localId: $localId, system: $system) { upc commercialType project { ids { localId } } } }""" GET_PRODUCT_BY_RELEASE_ARTIST_ID_QUERY = """ query($releaseArtistId: String!) { getProductByReleaseArtistId(releaseArtistId: $releaseArtistId) { upc ids { localId } commercialType } }""" GET_TRACK_BY_ID_QUERY = """ query($localId: String!, $system: System!) { getTrackByLocalId(localId: $localId, system: $system) { isrc } }""" GET_TRACK_BY_TRACK_ARTIST_ID_QUERY = """ query($trackArtistId: String!) { getTrackByTrackArtistId(trackArtistId: $trackArtistId) { isrc ids { localId } } }""" class GraphQLConnector: """Connector to an upstream GraphQL service.""" def __init__(self, url, secret, shared_secret_name, credential_fetcher=lambda: {}): """Create an upstream server instance.""" self.secret = secret self.shared_secret_name = shared_secret_name self.credential_fetcher = credential_fetcher self.token = None self.tokenExpiry = time.time() self.endpoint = HTTPEndpoint( url, base_headers={ "x-token": self._get_token(), "apollographql-client-name": "lambda-switchboard-producer", }, ) def _get_token(self): if self.token is None or self.tokenExpiry < time.time(): if self.secret is not None: secret = {"SecretString": self.secret, "VersionId": "unused"} else: credentials = self.credential_fetcher() secret = get_secret_value(self.shared_secret_name, credentials) token_payload = { "systemName": "ORCHARD", "exp": time.time() + 7 * 24 * 60 * 60, "iat": time.time(), "secretId": secret["VersionId"], } self.token = jwt.encode(token_payload, secret["SecretString"]) self.tokenExpiry = time.time() + 24 * 60 * 60 return self.token def _execute(self, query, data, correlation_id): """Execute a call to a GraphQL endpoint.""" return self.endpoint( query, data, extra_headers={ "Correlation-Id": correlation_id, "x-token": self._get_token(), }, ) class OrchardClient(GraphQLConnector): """Connector for graphql-switchboard.""" def __init__(self, url, secret, secret_name, credential_fetcher=lambda: {}): """Init OrchardClient.""" super().__init__(url, secret, secret_name, credential_fetcher) def lookup_product_by_id(self, release_id): """Lookup product using getProductByLocalId.""" data = {"localId": str(release_id), "system": "ORCHARD"} response = self._execute(GET_PRODUCT_BY_ID_QUERY, data, "") if response.get("errors"): return None result = response["data"]["getProductByLocalId"] if not result: return None return release_id, result def lookup_product_by_release_artist_id(self, release_artist_id): """Lookup product using getProductByReleaseArtistId.""" data = {"releaseArtistId": str(release_artist_id)} response = self._execute(GET_PRODUCT_BY_RELEASE_ARTIST_ID_QUERY, data, "") if response.get("errors"): return None result = response["data"]["getProductByReleaseArtistId"] if not result: return None return result["ids"][0]["localId"], result def lookup_track_by_id(self, track_id): """Lookup track using getTrackByLocalId.""" data = {"localId": str(track_id), "system": "ORCHARD"} response = self._execute(GET_TRACK_BY_ID_QUERY, data, "") if response.get("errors"): return None result = response["data"]["getTrackByLocalId"] if not result: return None return track_id, result def lookup_track_by_track_artist_id(self, track_artist_id): """Lookup track using getTrackByTrackArtistId.""" data = { "trackArtistId": str(track_artist_id), } response = self._execute(GET_TRACK_BY_TRACK_ARTIST_ID_QUERY, data, "") if response.get("errors"): return None result = response["data"]["getTrackByTrackArtistId"] if not result: return None return result["ids"][0]["localId"], result def get_orchard_client(): """Get Orchard client.""" return OrchardClient( config.GRAPHQL_URL, config.GRAPHQL_SECRET, config.GRAPHQL_SECRET_NAME )