CoolFace
Apppublic

ghstedpixel/BRIE

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes
app.py377 linesDownload Raw Back to root
1import os2import pickle3import subprocess4import sqlite35from datetime import datetime6 7import duckdb8import gradio as gr9import numpy as np10import pandas as pd11import pandas_ta as ta12import plotly.graph_objects as go13import yfinance as yf14from xgboost import XGBRegressor15from huggingface_hub import HfApi, hf_hub_download16from dotenv import load_dotenv17 18# --- HEALTH AGENT BLOCK ---19os.environ["PYTHONUNBUFFERED"] = "1"20 21def run_agent():22    print("๐Ÿฉบ Starting Health Agent...")23    result = subprocess.run(24        ["python", "health_agent.py"],25        cwd=".",26        capture_output=False,27        text=True28    )29    print(f"๐Ÿฉบ Agent complete: {result.returncode}")30 31run_agent()32print("๐Ÿฉบ Health Agent launched & complete!")33 34load_dotenv()35 36# --- UNIFIED PATHS & CONSTANTS ---37WATCHLIST_DB = os.path.join(os.getcwd(), "watchlist_manager.db")38DB_PATH = "/tmp/brie_brain.db"39MODEL_PATH = "/tmp/brie_model.pkl"40 41MOMENTUM_THRESHOLD = 7542OVERBOUGHT_RSI = 7043 44FEATURES = [45    "rsi", "macd_h", "ema_20_dist", "ema_50_dist", "vwap_dist",46    "atr_pct", "bb_width", "vol_spike", "ret_1d", "ret_5d",47    "day_of_week", "sentiment_score", "gap_pct", "vix_level"48]49 50SECTOR_ETFS = ["XLF", "XLK", "XLE", "XLY", "XLP", "XLV", "XLI", "XLU", "XLB", "XLRE", "XLC"]51VOLATILITY_INDEX = "^VIX"52 53HF_TOKEN = os.getenv("HF_TOKEN")54HF_REPO_ID = os.getenv("HF_REPO_ID")55 56# --- DATABASE HELPERS ---57def get_current_watchlist():58    conn = sqlite3.connect(WATCHLIST_DB)59    cursor = conn.cursor()60    cursor.execute("CREATE TABLE IF NOT EXISTS watchlist (ticker TEXT PRIMARY KEY)")61    cursor.execute("SELECT ticker FROM watchlist ORDER BY ticker")62    data = cursor.fetchall()63    conn.close()64    return pd.DataFrame(data, columns=["Ticker"])65 66def add_ticker_to_db(ticker):67    if ticker:68        ticker = ticker.upper().strip()69        conn = sqlite3.connect(WATCHLIST_DB)70        cursor = conn.cursor()71        cursor.execute("INSERT OR IGNORE INTO watchlist (ticker) VALUES (?)", (ticker,))72        conn.commit()73        conn.close()74    return get_current_watchlist(), ""75 76def remove_ticker_from_db(ticker):77    if ticker:78        ticker = ticker.upper().strip()79        conn = sqlite3.connect(WATCHLIST_DB)80        cursor = conn.cursor()81        cursor.execute("DELETE FROM watchlist WHERE ticker = ?", (ticker,))82        conn.commit()83        conn.close()84    return get_current_watchlist(), ""85 86def trigger_discord():87    try:88        from brie_agent import run_alpaca_briefing89        watchlist_df = get_current_watchlist()90        current_list = watchlist_df["Ticker"].tolist()91        if not current_list:92            return "โš ๏ธ Watchlist is empty."93        return str(run_alpaca_briefing(current_list))94    except Exception as e:95        return f"โŒ Bridge Error: {str(e)}"96 97# --- CLOUD SYNC & DB INIT ---98def init_db():99    with duckdb.connect(DB_PATH) as con:100        feat_cols = ", ".join([f'"{f}" FLOAT' for f in FEATURES])101        con.execute(102            f"""103            CREATE TABLE IF NOT EXISTS feature_store (104                ticker TEXT,105                timestamp TIMESTAMP,106                {feat_cols},107                target_next_day_ret FLOAT,108                PRIMARY KEY (ticker, timestamp)109            )110            """111        )112        con.execute(113            """114            CREATE TABLE IF NOT EXISTS performance_logs (115                ticker TEXT,116                timestamp TIMESTAMP,117                predicted_ret FLOAT,118                entry_price FLOAT,119                actual_ret FLOAT,120                is_verified BOOLEAN DEFAULT FALSE121            )122            """123        )124 125def load_from_cloud():126    if not HF_TOKEN or not HF_REPO_ID:127        return128    try:129        for f in ["brie_model.pkl", "brie_brain.db"]:130            hf_hub_download(131                repo_id=HF_REPO_ID,132                filename=f,133                local_dir="/tmp",134                repo_type="dataset",135                token=HF_TOKEN136            )137        print("โœ… Cloud models/data loaded")138    except Exception as e:139        print(f"โš ๏ธ Cloud load failed: {e}")140 141def save_to_cloud():142    if not HF_TOKEN or not HF_REPO_ID:143        return "โš ๏ธ HF_TOKEN missing"144    api = HfApi()145    try:146        for path, filename in [(MODEL_PATH, "brie_model.pkl"), (DB_PATH, "brie_brain.db")]:147            if os.path.exists(path):148                api.upload_file(149                    path_or_fileobj=path,150                    path_in_repo=filename,151                    repo_id=HF_REPO_ID,152                    repo_type="dataset",153                    token=HF_TOKEN154                )155        return "โœ… Sync Success"156    except Exception as e:157        return f"โŒ Sync Failed: {str(e)[:30]}"158 159# --- ANALYTICS ENGINE ---160def fetch_history_yf(ticker, period="5y", interval="1d"):161    try:162        df = yf.download(163            ticker,164            period=period,165            interval=interval,166            auto_adjust=False,167            progress=False168        )169        if df.empty:170            return pd.DataFrame()171        if isinstance(df.columns, pd.MultiIndex):172            df.columns = df.columns.get_level_values(0)173        df.columns = [str(c).lower().replace(" ", "_") for c in df.columns]174        df.index = pd.to_datetime(df.index).tz_localize(None)175        return df[["open", "high", "low", "close", "volume"]]176    except:177        return pd.DataFrame()178 179def build_feature_set(df, sentiment_val=0.0, labeling=False, vix_df=None):180    if len(df) < 100:181        return pd.DataFrame()182    ind = pd.DataFrame(index=df.index)183    ind["rsi"] = ta.rsi(df["close"], length=14).fillna(50.0)184    macd = ta.macd(df["close"])185    ind["macd_h"] = macd.iloc[:, 1].fillna(0.0) if macd is not None else 0.0186    bbands = ta.bbands(df["close"], length=20, std=2)187    ind["bb_width"] = ((bbands.iloc[:, 2] - bbands.iloc[:, 0]) / bbands.iloc[:, 1]).fillna(0.0) if bbands is not None else 0.0188    for length in [20, 50]:189        ema = ta.ema(df["close"], length=length)190        ind[f"ema_{length}_dist"] = ((df["close"] - ema) / ema).fillna(0.0)191    vwap = ta.vwap(df["high"], df["low"], df["close"], df["volume"])192    ind["vwap_dist"] = ((df["close"] - vwap) / vwap).fillna(0.0)193    ind["atr_pct"] = (ta.atr(df["high"], df["low"], df["close"]) / df["close"]).fillna(0.0)194    ind["vol_spike"] = (df["volume"] / df["volume"].rolling(20).mean()).fillna(1.0)195    ind["ret_1d"] = df["close"].pct_change(1).fillna(0.0)196    ind["ret_5d"] = df["close"].pct_change(5).fillna(0.0)197    ind["day_of_week"] = df.index.dayofweek.astype(float)198    ind["sentiment_score"] = float(sentiment_val)199    ind["gap_pct"] = ((df["open"] - df["close"].shift(1)) / df["close"].shift(1)).fillna(0.0)200    if vix_df is not None and not vix_df.empty:201        ind["vix_level"] = vix_df["close"].reindex(df.index, method="ffill").fillna(20.0)202    else:203        ind["vix_level"] = pd.Series(20.0, index=df.index)204    if labeling:205        ind["target_next_day_ret"] = df["close"].pct_change(1).shift(-1)206    cols = FEATURES + (["target_next_day_ret"] if labeling else [])207    return ind[cols].dropna()208 209def calculate_momentum_score(pred_return, rsi, vol_spike):210    normalized_pred = np.clip((pred_return + 5) * 10, 0, 100)211    vol_bonus = 15 if vol_spike > 1.2 else 0212    rsi_penalty = (rsi - OVERBOUGHT_RSI) * 2 if rsi > OVERBOUGHT_RSI else 0213    return np.clip((normalized_pred * 0.7) + vol_bonus - rsi_penalty, 0, 100)214 215def log_prediction_to_db(ticker, pred_val, entry_price):216    init_db()217    with duckdb.connect(DB_PATH) as con:218        con.execute(219            "INSERT INTO performance_logs (ticker, timestamp, predicted_ret, entry_price) VALUES (?, ?, ?, ?)",220            (ticker, datetime.now(), pred_val, entry_price)221        )222 223def update_performance_ledger():224    init_db()225    with duckdb.connect(DB_PATH) as con:226        pending = con.execute(227            "SELECT * FROM performance_logs WHERE is_verified = FALSE AND timestamp < (CURRENT_TIMESTAMP - INTERVAL '1 DAY')"228        ).df()229        for _, row in pending.iterrows():230            current_bars = fetch_history_yf(row["ticker"], period="1d")231            if not current_bars.empty:232                actual_ret = ((current_bars["close"].iloc[-1] - row["entry_price"]) / row["entry_price"]) * 100233                con.execute(234                    "UPDATE performance_logs SET actual_ret = ?, is_verified = TRUE WHERE ticker = ? AND timestamp = ?",235                    (actual_ret, row["ticker"], row["timestamp"])236                )237    return "โœ… Ledger Updated"238 239def get_performance_dashboard():240    update_performance_ledger()241    init_db()242    with duckdb.connect(DB_PATH) as con:243        df = con.execute("SELECT * FROM performance_logs WHERE is_verified = TRUE").df()244    if df.empty:245        return "### โš ๏ธ No verified data yet. (Wait 24h)", None246    df["error"] = df["predicted_ret"] - df["actual_ret"]247    win_rate = (np.sign(df["predicted_ret"]) == np.sign(df["actual_ret"])).mean() * 100248    fig = go.Figure()249    fig.add_trace(go.Scatter(x=df["predicted_ret"], y=df["actual_ret"], mode="markers", name="Predictions", marker=dict(color="#FFC107")))250    fig.update_layout(template="plotly_dark", title="Prediction vs. Actual", xaxis_title="AI Forecast %", yaxis_title="Actual Result %")251    stats = f"### Directional Win Rate: {win_rate:.1f}% | N={len(df)}"252    return stats, fig253 254def run_incremental_training(batch_str):255    tickers = [t.strip().upper() for t in batch_str.split(",") if t.strip()]256    vix_data = fetch_history_yf(VOLATILITY_INDEX)257    yield "๐Ÿ“ฅ Syncing Data..."258    all_data = []259    for t in tickers:260        bars = fetch_history_yf(t)261        if bars.empty:262            continue263        f_set = build_feature_set(bars, vix_df=vix_data, labeling=True)264        if f_set.empty:265            continue266        f_set["ticker"] = t267        f_set = f_set.reset_index().rename(columns={"Date": "timestamp", "index": "timestamp"})268        all_data.append(f_set)269    if not all_data:270        yield "โŒ No data found."271        return272    df = pd.concat(all_data, ignore_index=True)273    init_db()274    with duckdb.connect(DB_PATH) as con:275        con.register("df_tmp", df)276        con.execute("INSERT OR IGNORE INTO feature_store SELECT ticker, timestamp, " + ", ".join([f'"{f}"' for f in FEATURES]) + ", target_next_day_ret FROM df_tmp")277        data = con.execute("SELECT * FROM feature_store").df()278    yield "๐Ÿง  Training XGBoost..."279    model = XGBRegressor(n_estimators=200, max_depth=7, learning_rate=0.02)280    model.fit(data[FEATURES], data["target_next_day_ret"])281    with open(MODEL_PATH, "wb") as f:282        pickle.dump(model, f)283    yield f"โœ… Refreshed! {save_to_cloud()}"284 285def analyze_asset(ticker):286    ticker = ticker.upper().strip()287    bars = fetch_history_yf(ticker)288    if bars.empty:289        return "โŒ No data", None, "", "", ""290    last_price = bars["close"].iloc[-1]291    pred_text = "โš ๏ธ Model not trained"292    if os.path.exists(MODEL_PATH):293        with open(MODEL_PATH, "rb") as f:294            reg = pickle.load(f)295        feat = build_feature_set(bars, vix_df=fetch_history_yf(VOLATILITY_INDEX, period="1mo"))296        if not feat.empty:297            latest = feat.tail(1)298            pred_val = float(reg.predict(latest[FEATURES])[0] * 100)299            log_prediction_to_db(ticker, pred_val, last_price)300            m_score = calculate_momentum_score(pred_val, float(latest["rsi"].iloc[0]), float(latest["vol_spike"].iloc[0]))301            if m_score > MOMENTUM_THRESHOLD:302                signal, color = "๐Ÿ”ฅ STRONG", "#00FF00"303            elif m_score > 50:304                signal, color = "โš–๏ธ TREND", "#FFC107"305            else:306                signal, color = "โš ๏ธ WEAK", "#FF5252"307            pred_text = f"<div style='text-align: center; border: 1px solid {color}; padding: 10px;'><h2>{m_score:.1f}/100</h2><h3>{signal}</h3><p>Forecast: {pred_val:+.2f}%</p></div>"308    fig = go.Figure(data=[go.Candlestick(x=bars.index, open=bars["open"], high=bars["high"], low=bars["low"], close=bars["close"])])309    fig.update_layout(template="plotly_dark", xaxis_rangeslider_visible=False, height=450)310    return f"### {ticker} Analysis", fig, pred_text, f"${last_price:.2f}", ""311 312def get_sector_analysis():313    if not os.path.exists(MODEL_PATH):314        return "โŒ Train First", None, ""315    with open(MODEL_PATH, "rb") as f:316        reg = pickle.load(f)317    vix = fetch_history_yf(VOLATILITY_INDEX)318    rows = []319    for t in SECTOR_ETFS:320        bars = fetch_history_yf(t, period="2y")321        feat = build_feature_set(bars, vix_df=vix)322        if feat.empty:323            continue324        rows.append({"Sector": t, "Return": reg.predict(feat.tail(1)[FEATURES])[0] * 100})325    if not rows:326        return "โŒ No sector data", None, ""327    df_res = pd.DataFrame(rows).sort_values("Return", ascending=False)328    fig = go.Figure(go.Bar(x=df_res["Sector"], y=df_res["Return"], marker_color="#FFC107"))329    fig.update_layout(template="plotly_dark", title="Sector Sentiment", height=400)330    return f"### Strongest: {df_res['Sector'].iloc[0]}", fig, "Macro synced."331 332# --- GRADIO INTERFACE ---333with gr.Blocks() as brie:334    gr.Markdown("# ๐Ÿง€ Brie Ultra AI")335    with gr.Tabs():336        with gr.Tab("๐Ÿ” Scanner"):337            with gr.Row():338                ticker_in = gr.Textbox(label="Ticker", scale=4)339                scan_btn = gr.Button("Predict", variant="primary", scale=1)340            chart = gr.Plot()341            with gr.Row():342                res_txt, pred_out, price_lbl = gr.Markdown(), gr.HTML(), gr.Label(label="Last Close")343            scan_btn.click(analyze_asset, ticker_in, [res_txt, chart, pred_out, price_lbl])344        with gr.Tab("๐ŸŽฏ Accuracy Dashboard"):345            perf_btn = gr.Button("Refresh Accuracy Stats")346            perf_stats, perf_chart = gr.Markdown(), gr.Plot()347            perf_btn.click(get_performance_dashboard, None, [perf_stats, perf_chart])348        with gr.Tab("๐Ÿ“Š Sector Impact"):349            sector_btn = gr.Button("Analyze Macro")350            sector_top, sector_chart, sector_ai = gr.Markdown(), gr.Plot(), gr.Markdown()351            sector_btn.click(get_sector_analysis, None, [sector_top, sector_chart, sector_ai])352        with gr.Tab("โš™๏ธ Continuous Learning"):353            gr.Markdown("### Watchlist")354            with gr.Row():355                ticker_input = gr.Textbox(label="Ticker", scale=3)356                add_btn, remove_btn = gr.Button("โž• Add"), gr.Button("๐Ÿ—‘๏ธ Remove")357            watchlist_display = gr.Dataframe(headers=["Ticker"], value=get_current_watchlist())358            gr.Markdown("---")359            cal_btn = gr.Button("Learn & Sync", variant="stop")360            cal_out = gr.Markdown("Ready.")361            cal_btn.click(run_incremental_training, gr.Textbox(value=",".join(SECTOR_ETFS), visible=False), [cal_out])362            gr.Markdown("---")363            discord_btn = gr.Button("๐Ÿš€ Send Briefing", variant="primary")364            discord_status = gr.Label(value="Ready")365            add_btn.click(add_ticker_to_db, [ticker_input], [watchlist_display, ticker_input])366            remove_btn.click(remove_ticker_from_db, [ticker_input], [watchlist_display, ticker_input])367            discord_btn.click(trigger_discord, None, [discord_status])368 369if __name__ == "__main__":370    init_db()371    load_from_cloud()372    init_db()373    brie.queue().launch(374        theme=gr.themes.Default(primary_hue="amber"),375        server_name="0.0.0.0",376        server_port=int(os.getenv("PORT", 7860))377    )