summaryrefslogtreecommitdiff
path: root/plugins/warp/ui/shared/fetcher.cpp
blob: daa3fece5f5093aa918ca5812c395e2cb7497340 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
#include "fetcher.h"

#include <chrono>

WarpFetcher::WarpFetcher()
{
	m_logger = new BinaryNinja::Logger("WARP Fetcher");
}

void WarpFetcher::AddPendingFunction(const FunctionRef& func)
{
	std::lock_guard<std::mutex> lock(m_requestMutex);
	const auto guid = Warp::GetAnalysisFunctionGUID(*func);
	if (!guid.has_value() || m_processedGuids.contains(*guid))
		return;
	m_pendingRequests.push_back(func);
}

std::vector<FunctionRef> WarpFetcher::FlushPendingFunctions()
{
	std::lock_guard<std::mutex> lock(m_requestMutex);
	std::vector<FunctionRef> requests = std::move(m_pendingRequests);
	m_pendingRequests.clear();
	return requests;
}

void WarpFetcher::ExecuteCompletionCallback()
{
	std::vector<std::pair<CallbackId, CompletionCallback>> callbacks;
	{
		std::lock_guard<std::mutex> lock(m_requestMutex);
		callbacks.insert(callbacks.end(), m_completionCallbacks.begin(), m_completionCallbacks.end());
	}

	std::vector<CallbackId> toRemove = {};
	for (auto& [id, cb] : callbacks)
		if (cb() == RemoveCallback)
			toRemove.push_back(id);

	std::lock_guard<std::mutex> lock(m_requestMutex);
	for (auto id : toRemove)
		m_completionCallbacks.erase(id);
}

std::shared_ptr<WarpFetcher> WarpFetcher::Global()
{
	static auto global = std::make_shared<WarpFetcher>();
	return global;
}

void WarpFetcher::FetchPendingFunctions(const std::vector<Warp::SourceTag>& allowedTags)
{
	m_requestInProgress = true;
	const auto requests = FlushPendingFunctions();
	if (requests.empty())
	{
		m_logger->LogDebug("No pending requests to fetch... skipping");
		m_requestInProgress = false;
		return;
	}

	const auto start_time = std::chrono::high_resolution_clock::now();

	// Because we must fetch for a single target we map the function guids to the associated platform to perform fetches
	// for each.
	std::map<PlatformRef, std::unordered_set<Warp::FunctionGUID>> platformMappedGuidSet;
	std::map<PlatformRef, std::unordered_set<Warp::ConstraintGUID>> platformMappedConstraintSet;
	for (const auto& func : requests)
	{
		const auto warpFunc = Warp::Function::Get(*func);
		if (!warpFunc)
			continue;
		auto platform = func->GetPlatform();
		platformMappedGuidSet[platform].insert(warpFunc->GetGUID());

		// We want to keep track of the guids so we can constrain the server response to only return functions with any
		// of them.
		const auto constraints = warpFunc->GetConstraints();
		std::vector<Warp::ConstraintGUID> constraintGuids;
		constraintGuids.reserve(constraints.size());
		for (const auto& constraint : constraints)
			constraintGuids.push_back(constraint.guid);
		platformMappedConstraintSet[platform].insert(constraintGuids.begin(), constraintGuids.end());
	}

	std::map<PlatformRef, std::vector<Warp::FunctionGUID>> platformMappedGuids;
	for (const auto& [platform, guids] : platformMappedGuidSet)
		platformMappedGuids[platform] = std::vector(guids.begin(), guids.end());

	// We keep them in the set above so we don't duplicate a bunch for functions with the same set of constraint guids.
	std::map<PlatformRef, std::vector<Warp::ConstraintGUID>> platformMappedConstraints;
	for (const auto& [platform, guids] : platformMappedConstraintSet)
		platformMappedConstraints[platform] = std::vector(guids.begin(), guids.end());

	for (const auto& [platform, guids] : platformMappedGuids)
	{
		m_logger->LogDebugF("Fetching {} functions for platform {}", guids.size(), platform->GetName());
		auto target = Warp::Target::FromPlatform(*platform);
		for (const auto& container : Warp::Container::All())
			container->FetchFunctions(*target, guids, allowedTags, platformMappedConstraints[platform]);

		std::lock_guard<std::mutex> lock(m_requestMutex);
		for (const auto& guid : guids)
			m_processedGuids.insert(guid);
	}

	m_requestInProgress = false;
	ExecuteCompletionCallback();
	const auto end_time = std::chrono::high_resolution_clock::now();
	const std::chrono::duration<double> elapsed_time = end_time - start_time;
	m_logger->LogDebug("Fetch batch took %f seconds", elapsed_time.count());
}

void WarpFetcher::ClearProcessed()
{
	std::lock_guard<std::mutex> lock(m_requestMutex);
	m_logger->LogInfoF("Clearing {} processed functions from cache...", m_processedGuids.size());
	m_processedGuids.clear();
}