From 7f8d6a8131e4abfcda60dbc1fab965c4b044da8b Mon Sep 17 00:00:00 2001 From: caspar Date: Thu, 19 Mar 2026 14:43:28 -0400 Subject: [PATCH] Integrated FinBERT model. --- Backend.cs | 3 + Backend.csproj | 1 + Controller/NewsSentimentController.cs | 46 ++++ DatabaseHandler/NewsSentimentRecord.cs | 14 ++ DatabaseHandler/PostgresNewsSentimentStore.cs | 123 ++++++++++ Interface/INewsSentimentStore.cs | 19 ++ ServiceHandler/FinBertScoringService.cs | 220 ++++++++++++++++++ ServiceHandler/NewsSentimentService.cs | 101 ++++++++ 8 files changed, 527 insertions(+) create mode 100644 Controller/NewsSentimentController.cs create mode 100644 DatabaseHandler/NewsSentimentRecord.cs create mode 100644 DatabaseHandler/PostgresNewsSentimentStore.cs create mode 100644 Interface/INewsSentimentStore.cs create mode 100644 ServiceHandler/FinBertScoringService.cs create mode 100644 ServiceHandler/NewsSentimentService.cs diff --git a/Backend.cs b/Backend.cs index 451c7f1..68002eb 100644 --- a/Backend.cs +++ b/Backend.cs @@ -12,6 +12,9 @@ builder.Services.AddSwaggerGen(); builder.Services.AddSingleton(); builder.Services.AddSingleton(); builder.Services.AddSingleton(); +builder.Services.AddSingleton(); +builder.Services.AddSingleton(); +builder.Services.AddSingleton(); builder.Services.AddCors(options => { diff --git a/Backend.csproj b/Backend.csproj index fa2acbb..4e5c414 100644 --- a/Backend.csproj +++ b/Backend.csproj @@ -13,6 +13,7 @@ + diff --git a/Controller/NewsSentimentController.cs b/Controller/NewsSentimentController.cs new file mode 100644 index 0000000..9abd135 --- /dev/null +++ b/Controller/NewsSentimentController.cs @@ -0,0 +1,46 @@ +using Backend.ServiceHander; +using Microsoft.AspNetCore.Mvc; + +namespace Backend.Controller +{ + [ApiController] + [Route("api/news-sentiment")] + public sealed class NewsSentimentController : ControllerBase + { + private readonly NewsSentimentService _service; + + public NewsSentimentController(NewsSentimentService service) + { + _service = service; + } + + [HttpGet] + public async Task Analyze( + [FromQuery] string keyword, + CancellationToken cancellationToken) + { + try + { + var results = await _service.AnalyzeAsync(keyword, cancellationToken); + return Ok(results.Select(r => new + { + source = r.Source, + keyword = r.Keyword, + publishedAt = r.PublishedAt, + title = r.Title, + score = r.Score, + sentimentLabel = r.SentimentLabel, + url = r.Url + })); + } + catch (ArgumentException ex) + { + return BadRequest(ex.Message); + } + catch (Exception ex) + { + return StatusCode(500, new { error = ex.ToString() }); + } + } + } +} diff --git a/DatabaseHandler/NewsSentimentRecord.cs b/DatabaseHandler/NewsSentimentRecord.cs new file mode 100644 index 0000000..5e272ad --- /dev/null +++ b/DatabaseHandler/NewsSentimentRecord.cs @@ -0,0 +1,14 @@ +namespace Backend.DatabaseHandler +{ + public sealed class NewsSentimentRecord + { + public string Source { get; init; } = string.Empty; + public string Keyword { get; init; } = string.Empty; + public DateTimeOffset PublishedAt { get; init; } + public string Title { get; init; } = string.Empty; + public double Score { get; init; } + public string SentimentLabel { get; init; } = string.Empty; + public string TitleHash { get; init; } = string.Empty; + public string Url { get; init; } = string.Empty; + } +} diff --git a/DatabaseHandler/PostgresNewsSentimentStore.cs b/DatabaseHandler/PostgresNewsSentimentStore.cs new file mode 100644 index 0000000..5f38d68 --- /dev/null +++ b/DatabaseHandler/PostgresNewsSentimentStore.cs @@ -0,0 +1,123 @@ +using Backend.Interface; +using Npgsql; + +namespace Backend.DatabaseHandler +{ + public sealed class PostgresNewsSentimentStore : INewsSentimentStore + { + private readonly string _connectionString; + + public PostgresNewsSentimentStore(IConfiguration configuration) + { + _connectionString = configuration.GetConnectionString("MarketDataDb") + ?? throw new InvalidOperationException("Missing connection string: MarketDataDb"); + } + + public async Task> GetByKeywordAsync( + string keyword, + CancellationToken cancellationToken = default) + { + var results = new List(); + + await using var conn = new NpgsqlConnection(_connectionString); + await conn.OpenAsync(cancellationToken); + + const string sql = """ + SELECT source, keyword, published_at, title, score, + sentiment_label, title_hash, url + FROM news_sentiment_cache + WHERE keyword = @keyword + ORDER BY published_at DESC; + """; + + await using var cmd = new NpgsqlCommand(sql, conn); + cmd.Parameters.AddWithValue("keyword", keyword); + + await using var reader = await cmd.ExecuteReaderAsync(cancellationToken); + while (await reader.ReadAsync(cancellationToken)) + { + results.Add(new NewsSentimentRecord + { + Source = reader.GetString(0), + Keyword = reader.GetString(1), + PublishedAt = reader.GetFieldValue(2), + Title = reader.GetString(3), + Score = reader.GetDouble(4), + SentimentLabel = reader.GetString(5), + TitleHash = reader.GetString(6).TrimEnd(), + Url = reader.GetString(7) + }); + } + + return results; + } + + public async Task InsertAsync( + IEnumerable records, + CancellationToken cancellationToken = default) + { + var recordList = records.ToList(); + if (recordList.Count == 0) + return; + + await using var conn = new NpgsqlConnection(_connectionString); + await conn.OpenAsync(cancellationToken); + + await using var tx = await conn.BeginTransactionAsync(cancellationToken); + + const string sql = """ + INSERT INTO news_sentiment_cache + (source, keyword, published_at, title, score, + sentiment_label, title_hash, url) + VALUES + (@source, @keyword, @published_at, @title, @score, + @sentiment_label, @title_hash, @url) + ON CONFLICT (title_hash) DO NOTHING; + """; + + foreach (var record in recordList) + { + await using var cmd = new NpgsqlCommand(sql, conn, tx); + cmd.Parameters.AddWithValue("source", record.Source); + cmd.Parameters.AddWithValue("keyword", record.Keyword); + cmd.Parameters.AddWithValue("published_at", record.PublishedAt); + cmd.Parameters.AddWithValue("title", record.Title); + cmd.Parameters.AddWithValue("score", record.Score); + cmd.Parameters.AddWithValue("sentiment_label", record.SentimentLabel); + cmd.Parameters.AddWithValue("title_hash", record.TitleHash); + cmd.Parameters.AddWithValue("url", record.Url); + + await cmd.ExecuteNonQueryAsync(cancellationToken); + } + + await tx.CommitAsync(cancellationToken); + } + + public async Task> GetExistingHashesAsync( + string keyword, + CancellationToken cancellationToken = default) + { + var hashes = new HashSet(); + + await using var conn = new NpgsqlConnection(_connectionString); + await conn.OpenAsync(cancellationToken); + + const string sql = """ + SELECT title_hash + FROM news_sentiment_cache + WHERE keyword = @keyword; + """; + + await using var cmd = new NpgsqlCommand(sql, conn); + cmd.Parameters.AddWithValue("keyword", keyword); + + await using var reader = await cmd.ExecuteReaderAsync(cancellationToken); + while (await reader.ReadAsync(cancellationToken)) + { + hashes.Add(reader.GetString(0).TrimEnd()); + } + + return hashes; + } + } +} diff --git a/Interface/INewsSentimentStore.cs b/Interface/INewsSentimentStore.cs new file mode 100644 index 0000000..f11c9fe --- /dev/null +++ b/Interface/INewsSentimentStore.cs @@ -0,0 +1,19 @@ +using Backend.DatabaseHandler; + +namespace Backend.Interface +{ + public interface INewsSentimentStore + { + Task> GetByKeywordAsync( + string keyword, + CancellationToken cancellationToken = default); + + Task InsertAsync( + IEnumerable records, + CancellationToken cancellationToken = default); + + Task> GetExistingHashesAsync( + string keyword, + CancellationToken cancellationToken = default); + } +} diff --git a/ServiceHandler/FinBertScoringService.cs b/ServiceHandler/FinBertScoringService.cs new file mode 100644 index 0000000..67e745e --- /dev/null +++ b/ServiceHandler/FinBertScoringService.cs @@ -0,0 +1,220 @@ +using Microsoft.ML.OnnxRuntime; +using Microsoft.ML.OnnxRuntime.Tensors; +using System.Globalization; +using System.Text; + +namespace Backend.ServiceHander +{ + public sealed class FinBertScoringService : IDisposable + { + private readonly InferenceSession _session; + private readonly Dictionary _vocab; + private readonly string _inputIdsName; + private readonly string _attentionMaskName; + private readonly string _tokenTypeIdsName; + private readonly string _outputName; + + private const int MaxLength = 512; + private const int ClsId = 101; + private const int SepId = 102; + private const int UnkId = 100; + + // Label mapping from config.json: 0=positive, 1=negative, 2=neutral + private static readonly string[] Labels = ["positive", "negative", "neutral"]; + + public FinBertScoringService(IWebHostEnvironment env) + { + string modelDir = Path.Combine(env.ContentRootPath, "AIModels", "finbert"); + string modelPath = Path.Combine(modelDir, "model.onnx"); + string vocabPath = Path.Combine(modelDir, "vocab.txt"); + + _vocab = LoadVocab(vocabPath); + _session = new InferenceSession(modelPath); + + var inputNames = _session.InputMetadata.Keys.ToList(); + _inputIdsName = inputNames.First(n => n.Contains("input_id", StringComparison.OrdinalIgnoreCase)); + _attentionMaskName = inputNames.First(n => n.Contains("attention_mask", StringComparison.OrdinalIgnoreCase)); + _tokenTypeIdsName = inputNames.First(n => n.Contains("token_type", StringComparison.OrdinalIgnoreCase)); + _outputName = _session.OutputMetadata.Keys.First(); + } + + public (double Score, string Label) Score(string text) + { + var inputIds = Tokenize(text); + int seqLen = inputIds.Count; + + var inputIdsTensor = new DenseTensor(new[] { 1, seqLen }); + var attentionMaskTensor = new DenseTensor(new[] { 1, seqLen }); + var tokenTypeIdsTensor = new DenseTensor(new[] { 1, seqLen }); + + for (int i = 0; i < seqLen; i++) + { + inputIdsTensor[0, i] = inputIds[i]; + attentionMaskTensor[0, i] = 1; + tokenTypeIdsTensor[0, i] = 0; + } + + var inputs = new List + { + NamedOnnxValue.CreateFromTensor(_inputIdsName, inputIdsTensor), + NamedOnnxValue.CreateFromTensor(_attentionMaskName, attentionMaskTensor), + NamedOnnxValue.CreateFromTensor(_tokenTypeIdsName, tokenTypeIdsTensor) + }; + + using var results = _session.Run(inputs); + var outputTensor = results.First().AsTensor(); + float[] logits = outputTensor.ToArray(); + + var probs = Softmax(logits); + + // Score: positive * 100 + neutral * 50 + negative * 0 + double score = Math.Round(probs[0] * 100 + probs[2] * 50, 2); + + int maxIdx = Array.IndexOf(probs, probs.Max()); + string label = Labels[maxIdx]; + + return (score, label); + } + + private List Tokenize(string text) + { + var ids = new List { ClsId }; + + text = NormalizeText(text); + var basicTokens = BasicTokenize(text); + + foreach (var token in basicTokens) + { + ids.AddRange(WordPieceTokenize(token)); + } + + ids.Add(SepId); + + if (ids.Count > MaxLength) + { + ids = ids.Take(MaxLength - 1).ToList(); + ids.Add(SepId); + } + + return ids; + } + + private static string NormalizeText(string text) + { + text = text.ToLowerInvariant(); + + // Strip accents: NFD normalize then remove combining marks + text = new string(text.Normalize(NormalizationForm.FormD) + .Where(c => CharUnicodeInfo.GetUnicodeCategory(c) != UnicodeCategory.NonSpacingMark) + .ToArray()); + + var sb = new StringBuilder(text.Length); + foreach (char c in text) + { + if (char.IsControl(c)) + sb.Append(' '); + else + sb.Append(c); + } + + return sb.ToString(); + } + + private static List BasicTokenize(string text) + { + var tokens = new List(); + var current = new StringBuilder(); + + foreach (char c in text) + { + if (char.IsWhiteSpace(c)) + { + if (current.Length > 0) + { + tokens.Add(current.ToString()); + current.Clear(); + } + } + else if (char.IsPunctuation(c) || char.IsSymbol(c)) + { + if (current.Length > 0) + { + tokens.Add(current.ToString()); + current.Clear(); + } + tokens.Add(c.ToString()); + } + else + { + current.Append(c); + } + } + + if (current.Length > 0) + tokens.Add(current.ToString()); + + return tokens; + } + + private List WordPieceTokenize(string token) + { + var ids = new List(); + int start = 0; + + while (start < token.Length) + { + int end = token.Length; + bool found = false; + + while (start < end) + { + string sub = start > 0 + ? "##" + token[start..end] + : token[start..end]; + + if (_vocab.TryGetValue(sub, out int id)) + { + ids.Add(id); + found = true; + start = end; + break; + } + + end--; + } + + if (!found) + { + ids.Add(UnkId); + start++; + } + } + + return ids; + } + + private static float[] Softmax(float[] logits) + { + float max = logits.Max(); + var exps = logits.Select(l => MathF.Exp(l - max)).ToArray(); + float sum = exps.Sum(); + return exps.Select(e => e / sum).ToArray(); + } + + private static Dictionary LoadVocab(string vocabPath) + { + var vocab = new Dictionary(); + int index = 0; + foreach (string line in File.ReadLines(vocabPath)) + { + vocab[line.TrimEnd()] = index++; + } + return vocab; + } + + public void Dispose() + { + _session.Dispose(); + } + } +} diff --git a/ServiceHandler/NewsSentimentService.cs b/ServiceHandler/NewsSentimentService.cs new file mode 100644 index 0000000..b071417 --- /dev/null +++ b/ServiceHandler/NewsSentimentService.cs @@ -0,0 +1,101 @@ +using Backend.DatabaseHandler; +using Backend.Interface; +using System.Security.Cryptography; +using System.Text; +using System.Xml.Linq; + +namespace Backend.ServiceHander +{ + public sealed class NewsSentimentService + { + private readonly INewsSentimentStore _store; + private readonly FinBertScoringService _scorer; + private static readonly HttpClient Http = new() + { + DefaultRequestHeaders = { { "User-Agent", "SoftTrader/1.0" } } + }; + + public NewsSentimentService(INewsSentimentStore store, FinBertScoringService scorer) + { + _store = store; + _scorer = scorer; + } + + public async Task> AnalyzeAsync( + string keyword, + CancellationToken cancellationToken = default) + { + keyword = keyword.Trim(); + if (string.IsNullOrEmpty(keyword)) + throw new ArgumentException("Keyword cannot be empty."); + + var articles = await FetchRssAsync(keyword, cancellationToken); + + var existingHashes = await _store.GetExistingHashesAsync(keyword, cancellationToken); + + var newRecords = new List(); + foreach (var (title, url, publishedAt) in articles) + { + string titleHash = ComputeHash(title); + if (existingHashes.Contains(titleHash)) + continue; + + var (score, label) = _scorer.Score(title); + + newRecords.Add(new NewsSentimentRecord + { + Source = "google_rss", + Keyword = keyword, + PublishedAt = publishedAt, + Title = title, + Score = score, + SentimentLabel = label, + TitleHash = titleHash, + Url = url + }); + } + + if (newRecords.Count > 0) + await _store.InsertAsync(newRecords, cancellationToken); + + return await _store.GetByKeywordAsync(keyword, cancellationToken); + } + + private static async Task> FetchRssAsync( + string keyword, + CancellationToken cancellationToken) + { + string encodedKeyword = Uri.EscapeDataString(keyword); + string url = $"https://news.google.com/rss/search?q={encodedKeyword}&hl=en-US&gl=US&ceid=US:en"; + + string xml = await Http.GetStringAsync(url, cancellationToken); + var doc = XDocument.Parse(xml); + + var items = new List<(string, string, DateTimeOffset)>(); + + foreach (var item in doc.Descendants("item")) + { + string? title = item.Element("title")?.Value; + string? link = item.Element("link")?.Value; + string? pubDate = item.Element("pubDate")?.Value; + + if (string.IsNullOrWhiteSpace(title) || string.IsNullOrWhiteSpace(link)) + continue; + + DateTimeOffset published = DateTimeOffset.TryParse(pubDate, out var dt) + ? dt + : DateTimeOffset.UtcNow; + + items.Add((title, link, published)); + } + + return items; + } + + private static string ComputeHash(string text) + { + byte[] bytes = SHA256.HashData(Encoding.UTF8.GetBytes(text)); + return Convert.ToHexStringLower(bytes); + } + } +}