{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 1,
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "TONY\n"
     ]
    }
   ],
   "source": [
    "# make sure sqlalchemy is >=2\n",
    "# pip install psycopg2-binary\n",
    "# pip install \"sqlalchemy>=2\"\n",
    "import os\n",
    "from collections import defaultdict, Counter\n",
    "import json\n",
    "from urllib.parse import quote\n",
    "import time\n",
    "import pandas as pd\n",
    "import requests\n",
    "\n",
    "\n",
    "import boto3\n",
    "import pandas as pd\n",
    "import sqlalchemy\n",
    "import tqdm\n",
    "from botocore.exceptions import ClientError\n",
    "from snowflake.snowpark.functions import col\n",
    "\n",
    "\n",
    "# setup some pandas display stuff\n",
    "pd.set_option(\"display.max_rows\", 500)\n",
    "pd.set_option(\"display.max_columns\", 500)\n",
    "pd.set_option(\"display.width\", 1000)\n",
    "\n",
    "\n",
    "def get_secret():\n",
    "    secret_name = \"app-user-main-db-secret\"\n",
    "    region_name = \"us-east-2\"\n",
    "    # Create a Secrets Manager client\n",
    "    session = boto3.session.Session()\n",
    "    client = session.client(service_name=\"secretsmanager\", region_name=region_name)\n",
    "    try:\n",
    "        get_secret_value_response = client.get_secret_value(SecretId=secret_name)\n",
    "    except ClientError as e:\n",
    "        raise e\n",
    "    secret = get_secret_value_response[\"SecretString\"]\n",
    "    return json.loads(secret)\n",
    "\n",
    "\n",
    "my_secrets = get_secret()\n",
    "\n",
    "# alternative...\n",
    "engine = sqlalchemy.create_engine(\n",
    "    \"postgresql://suno:%s@suno-main-postgres-prod-analytics.cnfvffydbwvc.us-east-2.rds.amazonaws.com/suno_main\"\n",
    "    % quote(my_secrets[\"password\"]),\n",
    ")\n",
    "\n",
    "\n",
    "home_dir = os.path.expanduser(\"~\")\n",
    "snow_password_path = os.path.join(home_dir, \".aws\", \"snow_pw.txt\")\n",
    "if os.path.exists(snow_password_path):\n",
    "    # !pip install snowflake\n",
    "    from snowflake.core import Root\n",
    "    from snowflake.snowpark import Session\n",
    "\n",
    "    with open(snow_password_path, \"r\") as fp:\n",
    "        fp_lines = fp.readlines()\n",
    "        snow_password = fp_lines[0].strip()\n",
    "        snow_username = fp_lines[1].strip()\n",
    "    print(snow_username)\n",
    "    CONNECTION_PARAMETERS = {\n",
    "        \"account\": \"fu90569.us-east-2.aws\",\n",
    "        \"user\": snow_username,\n",
    "        \"private_key_file\": \"/Users/tonytong/.aws/snow_rsa_key.p8\",  # ask tony/jinhui for the key\n",
    "        \"role\": \"ACCOUNTADMIN\",\n",
    "        \"database\": \"SUNO_PROD\",\n",
    "        \"warehouse\": \"SUNO_PROD_LARGE\",\n",
    "        \"schema\": \"PROD\",\n",
    "    }"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 2,
   "metadata": {},
   "outputs": [],
   "source": [
    "with open(\"../src/app/groove_manifest_new.json\", \"r\") as f:\n",
    "    manifest = json.load(f)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 3,
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "(4674, 3)\n"
     ]
    }
   ],
   "source": [
    "flattened_clips = []\n",
    "for genre, clips in manifest.items():\n",
    "    # print(genre)\n",
    "    for clip in clips:\n",
    "        # print(clip)\n",
    "        clip.update({\"genre\": genre})\n",
    "        flattened_clips.append(clip)\n",
    "\n",
    "manifest_clips_df = pd.DataFrame(flattened_clips)\n",
    "print(manifest_clips_df.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 4,
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "PROD\n"
     ]
    }
   ],
   "source": [
    "if not os.path.exists(snow_password_path):\n",
    "    raise Exception(\"you are not authorized to access snowflake -- please setup\")\n",
    "\n",
    "snow_session = Session.builder.configs(CONNECTION_PARAMETERS).create()\n",
    "\n",
    "snow_root = Root(snow_session)\n",
    "snow_schema = snow_root.databases[\"SUNO_PROD\"].schemas[\"PROD\"]\n",
    "print(snow_schema.name)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 5,
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Total clips to query:  4674\n"
     ]
    },
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "  0%|          | 0/1 [00:00<?, ?it/s]"
     ]
    },
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Number of clip IDs in this chunk: 4674\n",
      "Length of the ID query string: 182285\n"
     ]
    },
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "100%|██████████| 1/1 [02:47<00:00, 167.67s/it]"
     ]
    },
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "1\n"
     ]
    },
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "\n"
     ]
    }
   ],
   "source": [
    "# Get clip IDs and query Snowflake in batches\n",
    "v4_clip_ids = list(str(s) for s in manifest_clips_df[\"clip_id\"].unique())\n",
    "snow_batch_size = 100_000\n",
    "snow_results = []\n",
    "print(\"Total clips to query: \", len(v4_clip_ids))\n",
    "for clip_ids_chunk in tqdm.tqdm(\n",
    "    [\n",
    "        v4_clip_ids[i : i + snow_batch_size]\n",
    "        for i in range(0, len(v4_clip_ids), snow_batch_size)\n",
    "    ]\n",
    "):\n",
    "    id_query_str = \",\".join(\"'\" + x + \"'\" for x in clip_ids_chunk)\n",
    "    print(f\"Number of clip IDs in this chunk: {len(clip_ids_chunk)}\")\n",
    "    print(f\"Length of the ID query string: {len(id_query_str)}\")\n",
    "\n",
    "    session_query = snow_session.sql(\n",
    "        f\"\"\" select *\n",
    "        from CLIP\n",
    "        where id in ({id_query_str})\n",
    "        \"\"\"\n",
    "    )\n",
    "    temp_df_snow_test = pd.DataFrame(session_query.collect())\n",
    "    snow_results.append(temp_df_snow_test)\n",
    "print(len(snow_results))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 6,
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Shape of df_snow_test:\n",
      "Rows: 4674\n",
      "Columns: 30\n"
     ]
    }
   ],
   "source": [
    "df_snow_test = pd.concat(snow_results)\n",
    "df_snow_test = df_snow_test.rename(columns=lambda x: x.lower())\n",
    "print(\"Shape of df_snow_test:\")\n",
    "print(f\"Rows: {df_snow_test.shape[0]}\")\n",
    "print(f\"Columns: {df_snow_test.shape[1]}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 7,
   "metadata": {},
   "outputs": [],
   "source": [
    "clip_to_lyrics = df_snow_test.groupby(\"id\")[\"prompt_text\"].first().to_dict()\n",
    "manifest_clips_df[\"prompt_text\"] = manifest_clips_df[\"clip_id\"].map(clip_to_lyrics)\n",
    "manifest_clips_df[\"prompt_text\"] = manifest_clips_df[\"prompt_text\"].apply(\n",
    "    lambda x: x.replace(\"{end}\", \"\\n[outro][end]\").strip()\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 8,
   "metadata": {},
   "outputs": [],
   "source": [
    "# TY pat.\n",
    "user_token = (\n",
    "    \"f5d11cd0bb704eb1b13320a0f6d8228f\"  # look up groovebot's token and put it here\n",
    ")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 9,
   "metadata": {},
   "outputs": [],
   "source": [
    "url = \"https://studio-api.prod.suno.com/api/generate/v2-web\"\n",
    "\n",
    "headers = {\n",
    "    \"Authorization\": f\"Bearer {user_token}\",\n",
    "    \"Content-Type\": \"application/json\",\n",
    "    \"Accept\": \"application/json\",\n",
    "    \"User-Agent\": \"Python/Requests\",\n",
    "    # \"x-suno-client\": \"ios\",\n",
    "}\n",
    "\n",
    "payload = {\n",
    "    \"generation_type\": \"TEXT\",\n",
    "    \"mv\": \"chirp-crow\",\n",
    "}"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 10,
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Total payloads:  4674\n"
     ]
    }
   ],
   "source": [
    "total_payloads = []\n",
    "for index, row in manifest_clips_df.iterrows():\n",
    "    new_payload = payload.copy()\n",
    "    new_payload[\"prompt\"] = row[\"prompt_text\"]\n",
    "    new_payload[\"tags\"] = row[\"genre\"]\n",
    "    new_payload[\"title\"] = (\n",
    "        row[\"title\"].capitalize() if row[\"title\"] else row[\"genre\"].capitalize()\n",
    "    )\n",
    "    total_payloads.append(new_payload)\n",
    "len(total_payloads)\n",
    "print(\"Total payloads: \", len(total_payloads))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 11,
   "metadata": {},
   "outputs": [
    {
     "ename": "NameError",
     "evalue": "name 'BREAK' is not defined",
     "output_type": "error",
     "traceback": [
      "\u001b[0;31m---------------------------------------------------------------------------\u001b[0m",
      "\u001b[0;31mNameError\u001b[0m                                 Traceback (most recent call last)",
      "Cell \u001b[0;32mIn[11], line 1\u001b[0m\n\u001b[0;32m----> 1\u001b[0m \u001b[43mBREAK\u001b[49m\n",
      "\u001b[0;31mNameError\u001b[0m: name 'BREAK' is not defined"
     ]
    }
   ],
   "source": [
    "BREAK"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 14,
   "metadata": {},
   "outputs": [],
   "source": [
    "# time.sleep(3600)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 15,
   "metadata": {},
   "outputs": [
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "Sending payloads: 100%|██████████| 4674/4674 [3:47:37<00:00,  2.92s/it]  "
     ]
    },
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "4674 0\n"
     ]
    },
    {
     "name": "stderr",
     "output_type": "stream",
     "text": [
      "\n"
     ]
    }
   ],
   "source": [
    "import copy\n",
    "import logging\n",
    "\n",
    "logger = logging.getLogger(\"clip_groove\")\n",
    "logger.setLevel(logging.INFO)\n",
    "\n",
    "\n",
    "def send_payload(payload, idx=None, max_retries=3, sleep_time=2.1):\n",
    "    \"\"\"\n",
    "    Send a single payload with retry logic.\n",
    "    Returns (success, response_data or error_info)\n",
    "    \"\"\"\n",
    "    for attempt in range(1, max_retries + 1):\n",
    "        try:\n",
    "            response = requests.post(\n",
    "                url=url,\n",
    "                headers=headers,\n",
    "                json=payload,\n",
    "                timeout=30,\n",
    "            )\n",
    "            response.raise_for_status()\n",
    "            return True, response.json()\n",
    "        except requests.exceptions.RequestException as e:\n",
    "            logger.warning(\n",
    "                f\"[{idx}] Attempt {attempt}/{max_retries} - Request error: {e} | Payload: {payload}\"\n",
    "            )\n",
    "            if (\n",
    "                hasattr(e, \"response\")\n",
    "                and e.response is not None\n",
    "                and hasattr(e.response, \"text\")\n",
    "            ):\n",
    "                logger.warning(f\"[{idx}] Response text: {e.response.text}\")\n",
    "            if attempt < max_retries:\n",
    "                time.sleep(sleep_time)\n",
    "            else:\n",
    "                return False, e\n",
    "        except Exception as e:\n",
    "            logger.error(f\"[{idx}] Unexpected error: {e}\", exc_info=True)\n",
    "            if attempt < max_retries:\n",
    "                time.sleep(sleep_time)\n",
    "            else:\n",
    "                return False, e\n",
    "    return False, None\n",
    "\n",
    "\n",
    "def process_payloads(payloads, sleep_time=2.1, max_retries=3):\n",
    "    \"\"\"\n",
    "    Process a list of payloads, sending each and collecting responses.\n",
    "    Returns (responses_dict, failed_payloads_list)\n",
    "    \"\"\"\n",
    "    responses = {}\n",
    "    failed = []\n",
    "\n",
    "    for i, payload in tqdm.tqdm(\n",
    "        enumerate(payloads), total=len(payloads), desc=\"Sending payloads\"\n",
    "    ):\n",
    "        success, result = send_payload(\n",
    "            payload, idx=i, max_retries=max_retries, sleep_time=sleep_time\n",
    "        )\n",
    "        if success:\n",
    "            responses[i] = result\n",
    "            time.sleep(sleep_time)\n",
    "        else:\n",
    "            failed.append((i, payload))\n",
    "    return responses, failed\n",
    "\n",
    "\n",
    "# First pass\n",
    "total_responses, failed_payloads = process_payloads(total_payloads)\n",
    "\n",
    "logger.info(\n",
    "    f\"First pass: {len(total_responses)} successes, {len(failed_payloads)} failures.\"\n",
    ")\n",
    "\n",
    "# Retry failed payloads once more\n",
    "if failed_payloads:\n",
    "    logger.info(f\"Retrying {len(failed_payloads)} failed payloads...\")\n",
    "    retry_indices, retry_payloads = (\n",
    "        zip(*failed_payloads) if failed_payloads else ([], [])\n",
    "    )\n",
    "    retry_responses, still_failed_payloads = process_payloads(\n",
    "        [p for _, p in failed_payloads]\n",
    "    )\n",
    "    # Map retry responses back to their original indices\n",
    "    for idx, resp in zip([i for i, _ in failed_payloads], retry_responses.values()):\n",
    "        total_responses[idx] = resp\n",
    "    failed_payloads = [\n",
    "        (i, p)\n",
    "        for (i, p), (success, _) in zip(\n",
    "            failed_payloads, [send_payload(p, idx=i) for i, p in failed_payloads]\n",
    "        )\n",
    "        if not success\n",
    "    ]\n",
    "    logger.info(\n",
    "        f\"After retry: {len(total_responses)} successes, {len(failed_payloads)} failures.\"\n",
    "    )\n",
    "\n",
    "print(len(total_responses), len(failed_payloads))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 16,
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "edec9e73-fb1c-40d7-bba7-a827ced6fe0d\n",
      "09747f5b-f18c-4350-889c-472829bca15b\n"
     ]
    }
   ],
   "source": [
    "for clip in total_responses[0][\"clips\"]:\n",
    "    print(clip[\"id\"])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 17,
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "Total clips generated:  939\n"
     ]
    }
   ],
   "source": [
    "flattened_responses = defaultdict(list)\n",
    "for response in total_responses.values():\n",
    "    for clip in response[\"clips\"]:\n",
    "        clip_result = {\n",
    "            \"id\": clip[\"id\"],\n",
    "            \"title\": clip[\"title\"],\n",
    "        }\n",
    "        flattened_responses[clip[\"metadata\"][\"tags\"]].append(clip_result)\n",
    "print(\"Total clips generated: \", len(flattened_responses))\n",
    "for genre, clips in flattened_responses.items():\n",
    "    flattened_responses[genre] = sorted(\n",
    "        flattened_responses[genre], key=lambda x: x[\"title\"]\n",
    "    )"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 18,
   "metadata": {},
   "outputs": [],
   "source": [
    "with open(\"../src/app/groove_manifest_crow.json\", \"w\") as f:\n",
    "    json.dump(flattened_responses, f, indent=4)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "suno_env",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "codemirror_mode": {
    "name": "ipython",
    "version": 3
   },
   "file_extension": ".py",
   "mimetype": "text/x-python",
   "name": "python",
   "nbconvert_exporter": "python",
   "pygments_lexer": "ipython3",
   "version": "3.10.12"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
