Java - Semaphore

1. 概述

Semaphore翻译成中文是信号量,是通过AQS实现的多线程工具类。
Semaphore信号量主要用于两个目的,一个是用于多个线程对多个共享资源的互斥使用,另一个是用于并发线程数的控制(限流)。

2. 简单使用

模拟 5 辆车停 3 个车位

public class SemaphoreTest {

    public static void main(String[] args) {
        Semaphore semaphore = new Semaphore(3);

        for (int i = 0; i < 5; i++) {
            new Thread(() -> {
                try {
                    semaphore.acquire();
                } catch (InterruptedException e) {
                    e.printStackTrace();
                }

                try {
                    System.out.println(Thread.currentThread().getName() + " start...");
                    TimeUnit.SECONDS.sleep(1);
                    System.out.println(Thread.currentThread().getName() + " end...");
                } catch (InterruptedException e) {
                    e.printStackTrace();
                } finally {
                    semaphore.release();
                }
            }).start();
        }
    }
}

3. 加锁解锁流程

Semaphore 有点像一个停车场,permits 就好像停车位数量,当线程获得了 permits 就像是获得了停车位,然后停车场显示空余车位减一
刚开始,permits(state)为 3,这时 5 个线程来获取资源

假设其中 Thread-1,Thread-2,Thread-4 cas 竞争成功,而 Thread-0 和 Thread-3 竞争失败,进入 AQS 队列 park 阻塞

这时 Thread-4 释放了 permits,状态如下

接下来 Thread-0 竞争成功,permits 再次设置为0,设置自己为 head 节点,断开原来的 head 节点,unpark 接下来的 Thread-3 节点,但由于 permits 是 0,因此 Thread-3 在尝试不成功后再次进入 park 状态

4. 源码

Semaphore 构造方法

// 从构造方法往下跟代码,最后permits就是设置为state
public Semaphore(int permits) {
    sync = new NonfairSync(permits);
}

static final class NonfairSync extends Sync {
    NonfairSync(int permits) {
        super(permits);
    }
}

abstract static class Sync extends AbstractQueuedSynchronizer {
    Sync(int permits) {
        setState(permits);
    }
}

protected final void setState(int newState) {
    state = newState;
}

Semaphore.acquire() 获取信号量方法

public void acquire() throws InterruptedException {
    sync.acquireSharedInterruptibly(1);
}

public final void acquireSharedInterruptibly(int arg)
        throws InterruptedException {
    if (Thread.interrupted())
        throw new InterruptedException();
    // 尝试获取信号量
    // 返回值大于0代表获取成功
    // 返回值小于0代表获取失败,会走下面的创建节点进入AQS队列的流程
    if (tryAcquireShared(arg) < 0)
        doAcquireSharedInterruptibly(arg);
}

// 模板方法具体看下面的实现,分为公平和非公平,我们这里跟非公平的代码
protected int tryAcquireShared(int arg) {
    throw new UnsupportedOperationException();
}

// 非公平的实现
static final class NonfairSync extends Sync {
    protected int tryAcquireShared(int acquires) {
        return nonfairTryAcquireShared(acquires);
    }
}

abstract static class Sync extends AbstractQueuedSynchronizer {
    final int nonfairTryAcquireShared(int acquires) {
        for (;;) {
            // 获取当前state的值(信号量)
            int available = getState();
            // 这里入参acquires为1,则当前state的值减去1
            int remaining = available - acquires;
            // 如果state-1小于0
            // 或者state-1大于等于0时,对state的值进行CAS为state-1成功
            if (remaining < 0 ||
                compareAndSetState(available, remaining))
                // 返回state-1后的信号量的值,由上层方法判断是否进入AQS队列流程
                return remaining;
        }
    }
}

// 判断获取信号量失败,走下面的创建节点进入AQS队列的流程
private void doAcquireSharedInterruptibly(int arg)
    throws InterruptedException {
    // 创建SHARED类型的Node节点,并加入AQS队列
    final Node node = addWaiter(Node.SHARED);
    boolean failed = true;
    try {
        for (;;) {
            // 获取当前节点的前置节点
            final Node p = node.predecessor();
            // 如果前置节点是头节点
            if (p == head) {
                // 尝试获取信号量
                int r = tryAcquireShared(arg);
                // 如果获取信号量返回值大于等于0
                if (r >= 0) {
                    // 将当前节点做为AQS队列的head节点
                    setHeadAndPropagate(node, r);
                    p.next = null; // help GC
                    failed = false;
                    return;
                }
            }
            // 如果获取信号量失败
            // shouldParkAfterFailedAcquire第一次循环将当前节点的前置节点waitStatus状态置为-1,然后返回false再次进入for (;;) 
            // shouldParkAfterFailedAcquire第二次循环会返回true,进入parkAndCheckInterrupt
            // parkAndCheckInterrupt会将当前节点进行park
            if (shouldParkAfterFailedAcquire(p, node) &&
                parkAndCheckInterrupt())
                throw new InterruptedException();
        }
    } finally {
        if (failed)
            cancelAcquire(node);
    }
}

Semaphore.release() 释放信号量方法

// 释放信号量
public void release() {
    sync.releaseShared(1);
}

public final boolean releaseShared(int arg) {
    // 如果尝试释放信号量成功,则走下面doReleaseShared的流程
    if (tryReleaseShared(arg)) {
        doReleaseShared();
        return true;
    }
    return false;
}

// 模板方法具体看下面的实现
protected boolean tryReleaseShared(int arg) {
    throw new UnsupportedOperationException();
}

protected final boolean tryReleaseShared(int releases) {
    for (;;) {
        // 获取当前state的值(信号量)
        int current = getState();
        // 这里入参releases为1,则当前state的值加上1
        int next = current + releases;
        if (next < current) // overflow
            throw new Error("Maximum permit count exceeded");
        // 如果对state的值进行CAS为state+1成功,则返回true,否则继续循环
        if (compareAndSetState(current, next))
            return true;
    }
}

// 尝试释放信号量成功走到这里
private void doReleaseShared() {
    for (;;) {
        // 拿到AQS队列的头节点,判断不为空不是最后一个
        Node h = head;
        if (h != null && h != tail) {
            // 获取头节点的waitStatus
            int ws = h.waitStatus;
            // 如果是-1则CAS改为0
            if (ws == Node.SIGNAL) {
                // 如果失败则进入for (;;) 进行重试
                if (!compareAndSetWaitStatus(h, Node.SIGNAL, 0))
                    continue;
                // 如果成功则进入unparkSuccessor流程,unpark头节点的后继节点(第二节点)
                // 这时候就需要跳转到线程的park位置
                unparkSuccessor(h);
            }
            else if (ws == 0 &&
                    !compareAndSetWaitStatus(h, 0, Node.PROPAGATE))
                continue;
        }
        if (h == head)
            break;
    }
}

private void doAcquireSharedInterruptibly(int arg)
    throws InterruptedException {
    final Node node = addWaiter(Node.SHARED);
    boolean failed = true;
    try {
        for (;;) {
            // 获取节点的前驱节点
            final Node p = node.predecessor();
            // 如果是头节点,说明自己是第二节点
            if (p == head) {
                // 尝试获得信号量
                int r = tryAcquireShared(arg);
                // 如果返回值大于等于0,说明获取信号量成功(r是表示剩余信号量的数值)
                if (r >= 0) {
                    // 将当前节点作为新的头节点,之前的头节点出队列
                    // 这里Semaphore还会继续尝试unpark后面的节点
                    // 但如果因为信号量为0,在尝试不成功后再次进入 park 状态
                    setHeadAndPropagate(node, r);
                    p.next = null; // help GC
                    failed = false;
                    return;
                }
            }
            if (shouldParkAfterFailedAcquire(p, node) &&
                // 被unpark的线程之前在这里被park了,unpark后回到这里,重新进入for (;;) 
                parkAndCheckInterrupt())
                throw new InterruptedException();
        }
    } finally {
        if (failed)
            cancelAcquire(node);
    }
}

private void setHeadAndPropagate(Node node, int propagate) {
    // 将当天节点作为头节点
    Node h = head;
    setHead(node);
    
    if (propagate > 0 || h == null || h.waitStatus < 0 ||
        (h = head) == null || h.waitStatus < 0) {
        // 这里会获取当前节点的后继节点,并且节点是SHARED的,则继续尝试去唤醒
        Node s = node.next;
        if (s == null || s.isShared())
            doReleaseShared();
    }
}
©著作权归作者所有,转载或内容合作请联系作者
【社区内容提示】社区部分内容疑似由AI辅助生成,浏览时请结合常识与多方信息审慎甄别。
平台声明:文章内容(如有图片或视频亦包括在内)由作者上传并发布,文章内容仅代表作者本人观点,简书系信息发布平台,仅提供信息存储服务。

相关阅读更多精彩内容

友情链接更多精彩内容