Files
whetstone_DSL/editor/src/VulnerabilityDatabase.h

396 lines
14 KiB
C
Raw Normal View History

#pragma once
// Step 190: Vulnerability knowledge base (OSV)
#include <string>
#include <vector>
#include <unordered_map>
#include <functional>
#include <filesystem>
#include <fstream>
#include <mutex>
#include <chrono>
#include <thread>
#include <algorithm>
#include <cctype>
#include <nlohmann/json.hpp>
struct VulnerabilityRecord {
std::string ecosystem;
std::string package;
std::string cveId;
std::string summary;
std::string severity;
std::vector<std::string> affectedVersions;
std::vector<std::string> fixedVersions;
std::vector<std::string> references;
struct VersionRange {
std::string introduced;
std::string fixed;
};
std::vector<VersionRange> ranges;
};
class VulnerabilityDatabase {
public:
using Fetcher = std::function<std::string(const std::string& ecosystem,
const std::string& package)>;
void setCacheDirectory(const std::string& path) {
std::scoped_lock lock(mu_);
cacheDir_ = path;
}
void setTTLHours(int hours) {
std::scoped_lock lock(mu_);
ttlHours_ = std::max(1, hours);
}
void setOfflineMode(bool offline) {
std::scoped_lock lock(mu_);
offlineMode_ = offline;
}
void setFetcher(Fetcher fetcher) {
std::scoped_lock lock(mu_);
fetcher_ = std::move(fetcher);
}
void trackPackage(const std::string& ecosystem,
const std::string& package) {
std::scoped_lock lock(mu_);
tracked_.push_back(makeKey(ecosystem, package));
}
void startBackgroundRefresh() {
std::thread([this]() { refreshTrackedPackages(); }).detach();
}
std::vector<VulnerabilityRecord> query(const std::string& ecosystem,
const std::string& package,
const std::string& version) {
ensureLoaded(ecosystem, package);
const std::string key = makeKey(ecosystem, package);
std::scoped_lock lock(mu_);
auto it = cache_.find(key);
if (it == cache_.end()) return {};
if (version.empty()) return it->second;
std::vector<VulnerabilityRecord> out;
for (const auto& record : it->second) {
if (isAffected(record, version)) out.push_back(record);
}
return out;
}
private:
std::string cacheDir_ = defaultCacheDir();
int ttlHours_ = 24;
bool offlineMode_ = false;
Fetcher fetcher_;
std::vector<std::string> tracked_;
std::unordered_map<std::string, std::vector<VulnerabilityRecord>> cache_;
std::mutex mu_;
static std::string defaultCacheDir() {
const char* home = std::getenv("USERPROFILE");
if (!home) home = std::getenv("HOME");
std::filesystem::path base = home ? home : ".";
return (base / ".whetstone" / "vuln_cache").string();
}
static std::string makeKey(const std::string& ecosystem,
const std::string& package) {
return ecosystem + ":" + package;
}
static std::string sanitizeFileName(const std::string& name) {
std::string out;
out.reserve(name.size());
for (char c : name) {
if (std::isalnum(static_cast<unsigned char>(c)) || c == '_' || c == '-') {
out.push_back(c);
} else if (c == ':' || c == '/' || c == '\\') {
out.push_back('_');
} else {
out.push_back('_');
}
}
return out;
}
std::filesystem::path cachePathFor(const std::string& ecosystem,
const std::string& package) const {
std::filesystem::path dir(cacheDir_);
std::string file = sanitizeFileName(makeKey(ecosystem, package)) + ".json";
return dir / file;
}
bool cacheIsStale(const std::filesystem::path& path) const {
if (!std::filesystem::exists(path)) return true;
auto last = std::filesystem::last_write_time(path);
auto now = std::filesystem::file_time_type::clock::now();
auto age = now - last;
auto hours = std::chrono::duration_cast<std::chrono::hours>(age).count();
return hours >= ttlHours_;
}
void ensureLoaded(const std::string& ecosystem,
const std::string& package) {
const std::string key = makeKey(ecosystem, package);
{
std::scoped_lock lock(mu_);
if (cache_.find(key) != cache_.end()) return;
}
refreshPackage(ecosystem, package, false);
}
void refreshTrackedPackages() {
std::vector<std::string> tracked;
{
std::scoped_lock lock(mu_);
tracked = tracked_;
}
for (const auto& key : tracked) {
auto pos = key.find(':');
if (pos == std::string::npos) continue;
refreshPackage(key.substr(0, pos), key.substr(pos + 1), true);
}
}
void refreshPackage(const std::string& ecosystem,
const std::string& package,
bool allowFetch) {
std::filesystem::path path = cachePathFor(ecosystem, package);
if (!allowFetch && !cacheIsStale(path)) {
loadCache(path, ecosystem, package);
return;
}
if (allowFetch && !offlineMode_ && fetcher_) {
std::string payload = fetcher_(ecosystem, package);
if (!payload.empty()) {
std::filesystem::create_directories(path.parent_path());
std::ofstream out(path);
if (out.is_open()) {
out << payload;
}
loadCache(path, ecosystem, package);
return;
}
}
loadCache(path, ecosystem, package);
}
void loadCache(const std::filesystem::path& path,
const std::string& ecosystem,
const std::string& package) {
std::ifstream in(path);
if (!in.is_open()) return;
nlohmann::json doc;
try {
in >> doc;
} catch (...) {
return;
}
auto records = parseOSV(doc);
std::scoped_lock lock(mu_);
cache_[makeKey(ecosystem, package)] = std::move(records);
}
static std::vector<VulnerabilityRecord> parseOSV(const nlohmann::json& doc) {
std::vector<VulnerabilityRecord> records;
if (!doc.is_object()) return records;
nlohmann::json vulns;
if (doc.contains("vulnerabilities")) {
vulns = doc.at("vulnerabilities");
} else if (doc.contains("id")) {
vulns = nlohmann::json::array({doc});
} else {
return records;
}
for (const auto& vuln : vulns) {
if (!vuln.is_object()) continue;
VulnerabilityRecord record;
record.summary = vuln.value("summary", "");
record.cveId = pickCVE(vuln);
record.severity = deriveSeverity(vuln);
record.references = parseReferences(vuln);
if (!vuln.contains("affected") || !vuln["affected"].is_array()) continue;
for (const auto& affected : vuln["affected"]) {
if (!affected.is_object()) continue;
if (!affected.contains("package")) continue;
const auto& pkg = affected["package"];
record.ecosystem = pkg.value("ecosystem", "");
record.package = pkg.value("name", "");
record.affectedVersions = parseVersions(affected);
record.fixedVersions = parseFixedVersions(affected);
record.ranges = parseRanges(affected);
records.push_back(record);
}
}
return records;
}
static std::vector<std::string> parseReferences(const nlohmann::json& vuln) {
std::vector<std::string> refs;
if (!vuln.contains("references")) return refs;
for (const auto& ref : vuln["references"]) {
if (!ref.is_object()) continue;
std::string url = ref.value("url", "");
if (!url.empty()) refs.push_back(url);
}
return refs;
}
static std::string pickCVE(const nlohmann::json& vuln) {
if (vuln.contains("aliases")) {
for (const auto& alias : vuln["aliases"]) {
if (alias.is_string()) {
std::string val = alias.get<std::string>();
if (val.rfind("CVE-", 0) == 0) return val;
}
}
}
return vuln.value("id", "");
}
static std::string deriveSeverity(const nlohmann::json& vuln) {
if (vuln.contains("severity") && vuln["severity"].is_array()) {
for (const auto& sev : vuln["severity"]) {
std::string score = sev.value("score", "");
double val = 0.0;
try {
val = std::stod(score);
} catch (...) {
val = 0.0;
}
if (val >= 9.0) return "Critical";
if (val >= 7.0) return "High";
if (val >= 4.0) return "Medium";
if (val > 0.0) return "Low";
}
}
if (vuln.contains("database_specific")) {
const auto& db = vuln["database_specific"];
if (db.contains("severity")) return db["severity"].get<std::string>();
}
return "Unknown";
}
static std::vector<std::string> parseVersions(const nlohmann::json& affected) {
std::vector<std::string> versions;
if (affected.contains("versions") && affected["versions"].is_array()) {
for (const auto& v : affected["versions"]) {
if (v.is_string()) versions.push_back(v.get<std::string>());
}
}
return versions;
}
static std::vector<std::string> parseFixedVersions(const nlohmann::json& affected) {
std::vector<std::string> versions;
if (!affected.contains("ranges") || !affected["ranges"].is_array()) return versions;
for (const auto& range : affected["ranges"]) {
if (!range.contains("events")) continue;
for (const auto& event : range["events"]) {
if (event.contains("fixed")) {
versions.push_back(event["fixed"].get<std::string>());
}
}
}
return versions;
}
static std::vector<VulnerabilityRecord::VersionRange>
parseRanges(const nlohmann::json& affected) {
std::vector<VulnerabilityRecord::VersionRange> ranges;
if (!affected.contains("ranges") || !affected["ranges"].is_array()) return ranges;
for (const auto& range : affected["ranges"]) {
if (!range.contains("events")) continue;
VulnerabilityRecord::VersionRange current;
bool hasOpen = false;
for (const auto& event : range["events"]) {
if (event.contains("introduced")) {
current.introduced = event["introduced"].get<std::string>();
current.fixed.clear();
hasOpen = true;
} else if (event.contains("fixed")) {
if (!hasOpen) current.introduced = "0";
current.fixed = event["fixed"].get<std::string>();
ranges.push_back(current);
hasOpen = false;
} else if (event.contains("last_affected")) {
if (!hasOpen) current.introduced = "0";
current.fixed = event["last_affected"].get<std::string>();
ranges.push_back(current);
hasOpen = false;
}
}
if (hasOpen) {
ranges.push_back(current);
}
}
return ranges;
}
static bool isAffected(const VulnerabilityRecord& record,
const std::string& version) {
for (const auto& v : record.affectedVersions) {
if (v == version) return true;
}
for (const auto& range : record.ranges) {
if (versionInRange(version, range.introduced, range.fixed)) return true;
}
return false;
}
static int compareVersions(const std::string& a, const std::string& b) {
auto split = [](const std::string& v) {
std::vector<std::string> parts;
std::string cur;
for (char c : v) {
if (std::isalnum(static_cast<unsigned char>(c)) || c == '.') {
cur.push_back(c);
} else if (!cur.empty()) {
parts.push_back(cur);
cur.clear();
}
}
if (!cur.empty()) parts.push_back(cur);
return parts;
};
auto pa = split(a);
auto pb = split(b);
size_t count = std::max(pa.size(), pb.size());
for (size_t i = 0; i < count; ++i) {
std::string sa = i < pa.size() ? pa[i] : "0";
std::string sb = i < pb.size() ? pb[i] : "0";
if (isNumber(sa) && isNumber(sb)) {
long long ia = std::stoll(sa);
long long ib = std::stoll(sb);
if (ia < ib) return -1;
if (ia > ib) return 1;
} else {
if (sa < sb) return -1;
if (sa > sb) return 1;
}
}
return 0;
}
static bool isNumber(const std::string& s) {
if (s.empty()) return false;
return std::all_of(s.begin(), s.end(),
[](unsigned char c) { return std::isdigit(c); });
}
static bool versionInRange(const std::string& version,
const std::string& introduced,
const std::string& fixed) {
if (!introduced.empty() && compareVersions(version, introduced) < 0) return false;
if (!fixed.empty() && compareVersions(version, fixed) >= 0) return false;
return true;
}
};