Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 7 additions & 3 deletions githubkit/cache/mem_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,16 +27,20 @@ def expire(self):

@override
def get(self, key: str) -> str | None:
self.expire()
return (item := self._cache.get(key, None)) and item.value
item = self._cache.get(key)
if item is None:
return None
if item.expire_at is not None and item.expire_at < datetime.now(timezone.utc):
self._cache.pop(key, None)
return None
return item.value

@override
async def aget(self, key: str) -> str | None:
return self.get(key)

@override
def set(self, key: str, value: str, ex: timedelta) -> None:
self.expire()
self._cache[key] = _Item(value, datetime.now(timezone.utc) + ex)

@override
Expand Down
16 changes: 11 additions & 5 deletions githubkit/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
from datetime import datetime, timedelta, timezone
import time
from types import TracebackType
from typing import TYPE_CHECKING, Any, Generic, TypeVar, cast, overload
from typing import TYPE_CHECKING, Any, Generic, TypeVar, overload

import anyio
import httpx
Expand Down Expand Up @@ -222,8 +222,11 @@ def __exit__(
exc_value: BaseException | None = None,
traceback: TracebackType | None = None,
):
cast(httpx.Client, self.__sync_client.get()).close()
self.__sync_client.set(None)
if client := self.__sync_client.get():
try:
client.close()
finally:
self.__sync_client.set(None)

# async context
async def __aenter__(self):
Expand All @@ -238,8 +241,11 @@ async def __aexit__(
exc_value: BaseException | None = None,
traceback: TracebackType | None = None,
):
await cast(httpx.AsyncClient, self.__async_client.get()).aclose()
self.__async_client.set(None)
if client := self.__async_client.get():
try:
await client.aclose()
finally:
self.__async_client.set(None)

def _get_client_defaults(self) -> dict[str, Any]:
"""Get default arguments for creating a httpx client."""
Expand Down
9 changes: 7 additions & 2 deletions githubkit/throttling.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,19 +37,24 @@ class LocalThrottler(BaseThrottler):

def __init__(self, max_concurrency: int) -> None:
self.max_concurrency = max_concurrency
self._lock = threading.Lock()
self._semaphore: threading.Semaphore | None = None
self._async_semaphore: anyio.Semaphore | None = None

@property
def semaphore(self) -> threading.Semaphore:
if self._semaphore is None:
self._semaphore = threading.Semaphore(self.max_concurrency)
with self._lock:
if self._semaphore is None:
self._semaphore = threading.Semaphore(self.max_concurrency)
return self._semaphore

@property
def async_semaphore(self) -> anyio.Semaphore:
if self._async_semaphore is None:
self._async_semaphore = anyio.Semaphore(self.max_concurrency)
with self._lock:
if self._async_semaphore is None:
self._async_semaphore = anyio.Semaphore(self.max_concurrency)
return self._async_semaphore

@override
Expand Down
52 changes: 51 additions & 1 deletion tests/test_unit_test/test_unit_test.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,17 @@
from datetime import timedelta
import json
from pathlib import Path
import threading
from typing import Any, TypeVar

from githubkit_schemas.latest.models import FullRepository
import httpx
import pytest

from githubkit import GitHub
from githubkit import GitHub, GitHubCore
from githubkit.cache.mem_cache import MemCache
from githubkit.response import Response
from githubkit.throttling import LocalThrottler
from githubkit.typing import UnsetType, URLTypes
from githubkit.utils import UNSET

Expand Down Expand Up @@ -76,3 +80,49 @@ async def test_async_mock():

repo = await target_async_func()
assert isinstance(repo, FullRepository)


def test_local_throttler_thread_safety():
throttler = LocalThrottler(max_concurrency=5)
threads = []
semaphores = []

def get_sem():
semaphores.append(throttler.semaphore)

for _ in range(20):
t = threading.Thread(target=get_sem)
threads.append(t)
t.start()

for t in threads:
t.join()

assert len(semaphores) == 20
first_sem = semaphores[0]
for sem in semaphores:
assert sem is first_sem


def test_mem_cache_passive_expiry():
cache = MemCache()
cache.set("key1", "val1", timedelta(milliseconds=1))
cache.set("key2", "val2", timedelta(hours=1))

import time

time.sleep(0.01)

assert cache.get("key1") is None
assert cache.get("key2") == "val2"
assert "key2" in cache._cache


def test_core_context_manager_safety():
gh = GitHubCore()
with gh:
with pytest.raises(RuntimeError):
gh.__enter__()

# Ensure no lingering client after error
assert gh._GitHubCore__sync_client.get() is None
Loading