def solution(citations):
    lo, hi = 0, len(citations) + 1
    while lo + 1 < hi:
        mid = (lo + hi) // 2
        cnt = 0
        for paper in citations:
            if paper >= mid:
                cnt += 1
        if cnt >= mid:
            lo = mid
        else:
            hi = mid
    return lo
    
print(solution([3, 0, 6, 1, 5]))
