diff --git a/prometheus_client/metrics.py b/prometheus_client/metrics.py index e3fe5323..dd20a33b 100644 --- a/prometheus_client/metrics.py +++ b/prometheus_client/metrics.py @@ -178,6 +178,11 @@ def labels(self: T, *labelvalues: object, **labelkwargs: object) -> T: if len(labelvalues) != len(self._labelnames): raise ValueError('Incorrect label count') str_labelvalues = tuple(str(l) for l in labelvalues) + try: + return self._metrics[str_labelvalues] + except KeyError: + pass + with self._lock: if str_labelvalues not in self._metrics: diff --git a/tests/test_core.py b/tests/test_core.py index 3aa19c24..a502cfdb 100644 --- a/tests/test_core.py +++ b/tests/test_core.py @@ -1,5 +1,6 @@ from concurrent.futures import ThreadPoolExecutor import os +from threading import Barrier, Event import time import unittest @@ -621,6 +622,58 @@ def test_child(self): self.two_labels.labels('x', 'y').inc(2) self.assertEqual(2, self.registry.get_sample_value('two_total', {'a': 'x', 'b': 'y'})) + def test_existing_child_during_other_child_creation(self): + started = Event() + release = Event() + + class SlowCounter(Counter): + def _metric_init(self): + if self._labelvalues == ('slow',): + started.set() + if not release.wait(timeout=10): + raise RuntimeError('Child creation was not released') + super()._metric_init() + + counter = SlowCounter('slow', 'help', ['l'], registry=self.registry) + existing = counter.labels('existing') + with ThreadPoolExecutor(max_workers=2) as pool: + creating = pool.submit(counter.labels, 'slow') + try: + self.assertTrue(started.wait(timeout=5)) + lookup = pool.submit(counter.labels, 'existing') + self.assertIs(existing, lookup.result(timeout=5)) + finally: + release.set() + self.assertIs(creating.result(timeout=5), counter.labels('slow')) + + def test_concurrent_child_creation_and_updates(self): + barrier = Barrier(8) + + def increment(): + barrier.wait(timeout=5) + child = self.counter.labels('shared') + for _ in range(100): + self.counter.labels(l='shared').inc() + return child + + with ThreadPoolExecutor(max_workers=8) as pool: + futures = [pool.submit(increment) for _ in range(8)] + children = [future.result(timeout=5) for future in futures] + + self.assertTrue(all(child is children[0] for child in children)) + self.assertEqual(800, self.registry.get_sample_value('c_total', {'l': 'shared'})) + + def test_recreate_child_after_removal(self): + for remove in (lambda: self.counter.remove('x'), self.counter.clear, + lambda: self.counter.remove_by_labels({'l': 'x'})): + original = self.counter.labels('x') + original.inc() + remove() + replacement = self.counter.labels('x') + self.assertIsNot(original, replacement) + self.assertIs(replacement, self.counter.labels(l='x')) + self.assertEqual(0, self.registry.get_sample_value('c_total', {'l': 'x'})) + def test_remove(self): self.counter.labels('x').inc() self.counter.labels('y').inc(2)