{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 18,
   "metadata": {},
   "outputs": [],
   "source": [
    "import pandas as pd\n",
    "import pandas as pd \n",
    "import numpy as np\n",
    "import matplotlib.pyplot as plt\n",
    "import seaborn as sns\n",
    "import math\n",
    "import datetime\n",
    "from string import ascii_letters\n",
    "from dateutil.relativedelta import relativedelta\n",
    "import math\n",
    "import pickle as pkl\n",
    "import swat\n",
    "import warnings\n",
    "import psycopg2\n",
    "import plotly.graph_objects as go\n",
    "from tqdm import tqdm\n",
    "import plotly\n",
    "warnings.filterwarnings('ignore')\n",
    "pd.set_option('display.max_columns', None)\n",
    "sns.set_style('darkgrid')\n",
    "%matplotlib inline"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 19,
   "metadata": {},
   "outputs": [],
   "source": [
    "path = '/Users/zhan009/Desktop/anomaly-detection/'\n",
    "stream_data = pd.read_csv(path + \"all_countries_3month_30s.csv\", index_col=0)\n",
    "stream_data = stream_data.sort_values(['primary_artist_name', 'track_name', 'track_isrc', 'min_report_date'])\n",
    "stream_data = stream_data.drop('track_album_artist', axis = 1)\n",
    "\n",
    "stream_data['cumsum_streams'] = stream_data.groupby(['primary_artist_name', 'track_name', 'track_isrc','country_code', 'min_report_date'])['streams'].cumsum(skipna=True)\n",
    "stream_data['cumsum_streams'] = np.log10(stream_data['cumsum_streams'])\n",
    "\n",
    "cumsum_table = stream_data.drop('streams', axis = 1) ### Table for cumulative sum\n",
    "pct_table = stream_data.drop('cumsum_streams', axis = 1) ### Table for percentage change"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "###### Stream Count"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 20,
   "metadata": {},
   "outputs": [],
   "source": [
    "def revert_data(dataset, columns):\n",
    "    dataset = pd.pivot_table(dataset, index=['primary_artist_name', 'track_isrc', 'track_name',\n",
    "                                             'country_code', 'min_report_date'], columns=['date_diff'], values=columns)\n",
    "    dataset = dataset.reset_index()\n",
    "    return dataset"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 21,
   "metadata": {},
   "outputs": [],
   "source": [
    "rt_cumsum = revert_data(cumsum_table, 'cumsum_streams')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 6,
   "metadata": {},
   "outputs": [],
   "source": [
    "# rt_cumsum.to_csv(\"cum_sum_stream_3_month.csv\")"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "###### Stream Percent Change"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 6,
   "metadata": {},
   "outputs": [],
   "source": [
    "percent_change = pct_table.copy()\n",
    "percent_change['streams'] = percent_change['streams'] + 1\n",
    "percent_change['pct_change'] = percent_change.groupby(['primary_artist_name', 'track_isrc', 'track_name', \n",
    "   \n",
    "                                                       'country_code', 'min_report_date'])['streams'].pct_change()\n",
    "rt_pct = revert_data(percent_change, 'pct_change')"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## Get Artist Quantile"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "- Example: Lil Nas X, or Camila Cebello"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "- First, get the 99 percentile for all the artist for 90 days. "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 22,
   "metadata": {},
   "outputs": [],
   "source": [
    "def get_artist_data(artist_name, dataset):\n",
    "    artist_data = dataset[dataset.primary_artist_name == artist_name]\n",
    "    return artist_data\n",
    "\n",
    "\n",
    "\n",
    "'''\n",
    "remove duplicate with same track name, track_isrc, but different min_report date\n",
    "only keep the earliest min_report_date\n",
    "'''\n",
    "\n",
    "def remove_duplicate(dataset):\n",
    "    artist_data = pd.DataFrame()\n",
    "    for i in set(dataset.track_isrc):\n",
    "        track_data = dataset[dataset.track_isrc == i]\n",
    "        date = track_data.groupby(['track_name', 'primary_artist_name', 'track_isrc', 'country_code'])['min_report_date'].min().values\n",
    "        track_data = track_data[track_data['min_report_date'] == date[0]]\n",
    "        artist_data = artist_data.append(track_data)\n",
    "    return artist_data"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "## The performace of Artists in different country (for one artist, or one track)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 23,
   "metadata": {},
   "outputs": [],
   "source": [
    "def calculate_quantile_country(country_code, artist_quantile, quantile_level):\n",
    "    artist_quantile = artist_quantile[artist_quantile.country_code != country_code]\n",
    "    all_artist_quantile = artist_quantile.quantile(quantile_level, numeric_only=True)\n",
    "    artist_quantile_dict = dict(zip(all_artist_quantile.index, all_artist_quantile.values))\n",
    "    return all_artist_quantile, artist_quantile_dict\n",
    "\n",
    "def get_result_country(track_isrc, country_code, dataset, quantile_data, label, start_day, quantile):\n",
    "    streams = dataset[range(start_day, 60)]\n",
    "    streams = streams.dropna(axis = 1)\n",
    "    if not streams.empty:\n",
    "        if streams.shape[1] > 10:\n",
    "#             plt.figure(figsize=(20, 10))\n",
    "#             print(\"Track name: {0}; {1} days (from Day {2} to Day {3})that surpass the {4}% quantile for all tracks\".format(track_info[track_isrc], \\\n",
    "#                                                                                                    streams.shape[1], \\\n",
    "#                                                                                                    streams.columns[0], \\\n",
    "#                                                                                                    streams.columns[-1], \\\n",
    "#                                                                                                    quantile*100))\n",
    "#             sns.lineplot(list(streams.columns)[start_day:], streams.values[0][start_day:], label = country_code)\n",
    "#             fig.add_trace(go.Scatter(x=list(streams.columns)[start_day:], \\\n",
    "#                                      y=streams.values[0][start_day:], \\\n",
    "#                                      name=track_info[track_isrc] +'-'+country_code,\\\n",
    "#                                      line=dict(color='firebrick', width=3),\n",
    "#                                      mode = 'lines', \n",
    "#                                      connectgaps=True))\n",
    "#             fig.add_trace(go.Scatter(x=list(range(start_day, 60)), y=quantile_data[start_day:], \\\n",
    "#                          name=track_info[track_isrc] +'-'+ \\\n",
    "#                              '{0}% Quantile for All Other Country Except {1}'.format(quantile*100, country_code), \\\n",
    "#                          line=dict(color='yellow', width=3), line_shape='spline', mode = 'lines',\n",
    "#                          connectgaps=True,))\n",
    "#             sns.lineplot(range(start_day, 91), quantile_data[start_day:], label = label,lw = 3, color = 'yellow', legend='full')  \n",
    "#             plt.title(\"Track name: {0} performance in {1}\".format(track_info[track_isrc], country_code))\n",
    "            return country_code\n",
    "#         else:\n",
    "# #             print(\"There is no track that the streams is surpass the {0}% quantile for all tracks\".format(quantile*100))\n",
    "#             continue\n",
    "            \n",
    "def quantile_com_country(x, cal_dict):\n",
    "    if x.dtype != np.number:\n",
    "        return x\n",
    "    else:\n",
    "        x = x.where(x>(cal_dict[x.name]))\n",
    "        return x"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 30,
   "metadata": {},
   "outputs": [],
   "source": [
    "artist_name = 'Camila Cabello'\n",
    "track_name = 'Havana'\n",
    "quantile = 0.995\n",
    "\n",
    "artist_data_percentile = get_artist_data(artist_name, rt_cumsum)\n",
    "# artist_data_percentile_pct = get_artist_data(artist_name, rt_pct)\n",
    "track_info = dict(zip(artist_data_percentile.track_isrc, artist_data_percentile.track_name))\n",
    "track_isrc = [i for i in track_info if track_info[i] == track_name][0]\n",
    "artist_data_percentile = remove_duplicate(artist_data_percentile)\n",
    "\n",
    "track_history = {}\n",
    "# fig = go.Figure()\n",
    "# for track_isrc in set(artist_data_percentile.track_isrc):\n",
    "song_data = artist_data_percentile[artist_data_percentile.track_isrc == track_isrc]\n",
    "US_record = song_data[song_data.country_code == 'US'].dropna(axis = 1)\n",
    "country_list = set()\n",
    "for country_code in set(song_data.country_code):\n",
    "    country_data = song_data[song_data.country_code == country_code]\n",
    "    artist_quantile_history, quantile_dict = calculate_quantile_country(country_code, song_data, quantile)\n",
    "    country_data = country_data.apply(quantile_com_country, cal_dict = quantile_dict)\n",
    "    if 'US' in set(song_data.country_code) and US_record.shape[1] > 5:\n",
    "        countries = get_result_country(track_isrc, country_code, country_data, artist_quantile_history, \\\n",
    "                                      '{0}% Quantile for All Artists'.format(quantile * 100), \\\n",
    "                                      0, quantile)\n",
    "        if countries:\n",
    "            country_list.add(countries)\n",
    "if country_list:\n",
    "    track_history[track_info[track_isrc]] = country_list\n",
    "\n",
    "for i in track_history:\n",
    "    countries = track_history[i]\n",
    "    if len(countries) > 1:\n",
    "        print(\"Track Name: {0}, Country Codes: {1}\".format(track_name, countries))\n",
    "\n",
    "# fig.update_layout(title='',\n",
    "#                    xaxis_title='Days',\n",
    "#                    yaxis_title='Streams')\n",
    "        \n",
    "# fig.show()  "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 31,
   "metadata": {},
   "outputs": [
    {
     "data": {
      "text/plain": [
       "{'Havana': {'US'}}"
      ]
     },
     "execution_count": 31,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "track_history"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python 3",
   "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.7.4"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
